diff --git a/README.md b/README.md index 0386526..7fca880 100644 --- a/README.md +++ b/README.md @@ -41,11 +41,33 @@ Provides straightforward in-memory conversation context management. This is the [[contexts]] module = "context-simple" name = "simple" -config = { - max_messages = 100 # Optional limit -} + +[contexts.config] +max_tool_result_bytes = 131072 # Optional override; default is 128 KiB ``` +### Tool-result text ingress cap + +Before admission, a tool message's string content (including JSON serialized +`ToolResult` dict/list output) is limited to `max_tool_result_bytes`, which +defaults to 131,072 UTF-8 bytes. Oversized content keeps a UTF-8-safe prefix +and one explicit retrieval marker. Raise this one explicit setting only for a +legitimate larger text result; there is no off switch in this version. + +For block content, the cap covers only direct text blocks. Image, audio, +unknown, and nested block data are not read or sliced. The original oversized +text is not retained in metadata, a side file, or the admitted transcript: +retrieve missing content using narrower read/query parameters; do not repeat +state-changing actions just to recover output. For ill-formed Python text +containing lone surrogates, byte accounting uses UTF-8 replacement; under-cap +content remains unchanged, while an oversized clipped prefix is valid UTF-8. + +The default is a finite observed baseline, not a universal tail guarantee: +509 outputs from 16 stock-main S1 captures had a 40,139-byte p99 and +87,301-byte maximum, with none above 128 KiB. 128 KiB is 3.27x that p99 and +1.5x that maximum, while still fitting the local 64k-window regression where +a 256 KiB cap would not. + ## Usage ```python @@ -67,7 +89,10 @@ Not suitable for: ## Compaction Strategy -The SimpleContextManager uses **ephemeral compaction** - `get_messages_for_request()` returns a compacted VIEW without modifying the internal message history. The full history is always preserved in memory. +The SimpleContextManager uses **ephemeral compaction** - +`get_messages_for_request()` returns a compacted VIEW without modifying the +admitted internal message history. Ingress-clipped tool text is irreversible +and is not retained in the canonical transcript; compaction remains view-only. Compaction triggers when token usage reaches the configured threshold (default: 92% of the **effective budget** -- which is derived from the provider, *not* from `max_tokens`; see [Where the compaction trigger comes from](#where-the-compaction-trigger-comes-from)): diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index 7c36e79..4435d13 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -7,7 +7,8 @@ • get_messages_for_request() returns compacted VIEW (new list) • get_messages() returns FULL history (for transcripts/session persistence) -This design ensures conversation history is never lost, even during compaction. +Compaction never loses admitted history. Oversized direct tool-result text is +irreversibly clipped before admission and is not retained in the transcript. For persistent storage across sessions, use context-persistent instead. Dynamic System Prompt Support: @@ -41,12 +42,25 @@ import logging from collections.abc import Awaitable, Callable from datetime import UTC, datetime +from sys import maxsize from typing import Any -from amplifier_core import ModuleCoordinator +from amplifier_core import ModuleCoordinator, TextBlock logger = logging.getLogger(__name__) +DEFAULT_MAX_TOOL_RESULT_BYTES = 128 * 1024 +_TOOL_RESULT_INGRESS_MARKER = ( + "[tool-result truncated at ingress: original_text_utf8_bytes={original_text_utf8_bytes}; " + "retrieve missing content using narrower read/query parameters; do not repeat " + "state-changing actions just to recover output.]" +) +_MIN_MAX_TOOL_RESULT_BYTES = len( + _TOOL_RESULT_INGRESS_MARKER.format(original_text_utf8_bytes=maxsize).encode( + "utf-8", errors="replace" + ) +) + # token_meter config values. "estimate" (default) preserves pre-existing # behavior exactly; "actual" lets a real llm:response measurement drive the # compaction trigger once one has arrived this session. See module @@ -120,6 +134,9 @@ async def mount(coordinator: ModuleCoordinator, config: dict[str, Any] | None = falling back to the estimator before then. An unrecognized value falls back to "estimate" with a logged warning rather than crashing mount(). See module docstring. + - max_tool_result_bytes: Maximum UTF-8 bytes admitted for direct + text in one tool result (default: 131,072). Oversized text is + irreversibly clipped at ingress with a retrieval marker. Returns: Cleanup callable that unregisters the token-meter hook (if one was @@ -152,6 +169,9 @@ async def mount(coordinator: ModuleCoordinator, config: dict[str, Any] | None = output_reserve_fraction=config.get("output_reserve_fraction", 0.5), token_meter=token_meter, hooks=getattr(coordinator, "hooks", None), + max_tool_result_bytes=config.get( + "max_tool_result_bytes", DEFAULT_MAX_TOOL_RESULT_BYTES + ), ) # Always register the meter listener when hooks are available, regardless @@ -189,7 +209,8 @@ class SimpleContextManager: Owns memory policy: orchestrators ask for messages via get_messages_for_request(), and this context manager decides how to fit them within limits. Compaction is - handled internally and ephemerally - the original history is always preserved. + handled internally and ephemerally - admitted history is always preserved. + Oversized direct tool-result text is clipped before it enters that history. Compaction Strategy (Progressive Interleaved): Triggered when usage >= compact_threshold (default 92%), target is target_usage (default 50%). @@ -228,6 +249,7 @@ def __init__( output_reserve_fraction: float = 0.5, token_meter: str = TOKEN_METER_ESTIMATE, hooks: Any = None, + max_tool_result_bytes: int = DEFAULT_MAX_TOOL_RESULT_BYTES, ): """ Initialize the context manager. @@ -257,7 +279,20 @@ def __init__( hooks: Optional hooks instance for emitting observability events and (always, when present) recording real usage for the token meter via `llm:response` -- see `_on_llm_response`. + max_tool_result_bytes: Maximum UTF-8 bytes admitted for direct + text in one tool result. The minimum allows the truncation + marker itself; default is 128 KiB. """ + if ( + isinstance(max_tool_result_bytes, bool) + or not isinstance(max_tool_result_bytes, int) + or max_tool_result_bytes < _MIN_MAX_TOOL_RESULT_BYTES + ): + raise ValueError( + "max_tool_result_bytes must be an integer at least " + f"{_MIN_MAX_TOOL_RESULT_BYTES} so the ingress truncation marker fits" + ) + self.messages: list[dict[str, Any]] = [] self.max_tokens = max_tokens self.compact_threshold = compact_threshold @@ -279,6 +314,7 @@ def __init__( token_meter = TOKEN_METER_ESTIMATE self.token_meter = token_meter self._hooks = hooks + self.max_tool_result_bytes = max_tool_result_bytes self._last_compaction_stats: dict[str, Any] | None = None # Real-usage token meter state (see _on_llm_response / # _measure_working_tokens). `_last_measured_prompt_tokens` holds the @@ -313,6 +349,143 @@ def __init__( # accumulated effect, not just the most recent escalation step. self._sticky_level: int = 0 + @staticmethod + def _encode_tool_result_text(text: str) -> bytes: + """Encode direct text, replacing lone surrogates for ingress accounting.""" + return text.encode("utf-8", errors="replace") + + @classmethod + def _utf8_safe_prefix(cls, text: str, max_bytes: int) -> str: + """Return a valid UTF-8 prefix, replacing lone surrogates when clipped.""" + return cls._encode_tool_result_text(text)[:max_bytes].decode( + "utf-8", errors="ignore" + ) + + @staticmethod + def _direct_text_block_text(block: Any) -> str | None: + """Return text only from the two supported direct text block shapes.""" + if isinstance(block, TextBlock): + return block.text + if ( + isinstance(block, dict) + and block.get("type") == "text" + and isinstance(block.get("text"), str) + ): + return block["text"] + return None + + @staticmethod + def _copy_text_block_with_text(block: Any, text: str) -> Any: + """Copy a recognized text block while preserving its other fields.""" + if isinstance(block, TextBlock): + return block.model_copy(update={"text": text}) + return {**block, "text": text} + + def _limit_tool_result_text_at_ingress( + self, message: dict[str, Any] + ) -> tuple[dict[str, Any], dict[str, Any] | None]: + """Bound direct tool-result text without mutating the caller's message. + + This intentionally handles only a string content value and direct + ``TextBlock``/``{"type": "text"}`` entries in a top-level content list. + Other blocks, nested structures, and unknown shapes remain untouched. + """ + if message.get("role") != "tool": + return message, None + + content = message.get("content") + content_kind: str | None = None + if isinstance(content, str): + original_text_utf8_bytes = len(self._encode_tool_result_text(content)) + content_kind = "string" + elif isinstance(content, list): + original_text_utf8_bytes = sum( + len(self._encode_tool_result_text(text)) + for block in content + if (text := self._direct_text_block_text(block)) is not None + ) + content_kind = "text_blocks" + else: + return message, None + + if original_text_utf8_bytes <= self.max_tool_result_bytes: + return message, None + + marker = _TOOL_RESULT_INGRESS_MARKER.format( + original_text_utf8_bytes=original_text_utf8_bytes + ) + marker_bytes = len(self._encode_tool_result_text(marker)) + remaining_text_bytes = self.max_tool_result_bytes - marker_bytes + + if isinstance(content, str): + limited_content = self._utf8_safe_prefix(content, remaining_text_bytes) + marker + else: + limited_content: list[Any] = [] + marker_added = False + for block in content: + text = self._direct_text_block_text(block) + if text is None: + limited_content.append(block) + continue + if marker_added: + continue + + prefix = self._utf8_safe_prefix(text, remaining_text_bytes) + prefix_bytes = len(self._encode_tool_result_text(prefix)) + if prefix_bytes == len(self._encode_tool_result_text(text)): + limited_content.append(self._copy_text_block_with_text(block, prefix)) + remaining_text_bytes -= prefix_bytes + continue + + limited_content.append( + self._copy_text_block_with_text(block, prefix + marker) + ) + marker_added = True + + stored_text_utf8_bytes = ( + len(self._encode_tool_result_text(limited_content)) + if isinstance(limited_content, str) + else sum( + len(self._encode_tool_result_text(text)) + for block in limited_content + if (text := self._direct_text_block_text(block)) is not None + ) + ) + event_data = { + "tool_name": message.get("name"), + "tool_call_id": message.get("tool_call_id"), + "original_text_utf8_bytes": original_text_utf8_bytes, + "stored_text_utf8_bytes": stored_text_utf8_bytes, + "max_tool_result_bytes": self.max_tool_result_bytes, + "content_kind": content_kind, + } + return {**message, "content": limited_content}, event_data + + async def _emit_tool_result_ingress_truncation(self, event_data: dict[str, Any]) -> None: + """Emit ingress observability without ever blocking tool-result admission.""" + if self._hooks is not None: + try: + await self._hooks.emit("context:tool_result_ingress_truncated", event_data) + return + except Exception: + emission_status = "event emission failed" + else: + emission_status = "no hooks are registered" + + logger.warning( + "context-simple: tool result truncated at ingress (%s; tool_name=%r, " + "tool_call_id=%r, original_text_utf8_bytes=%d, stored_text_utf8_bytes=%d, " + "max_tool_result_bytes=%d); retrieve missing content using narrower " + "read/query parameters; do not repeat state-changing actions just to " + "recover output.", + emission_status, + event_data["tool_name"], + event_data["tool_call_id"], + event_data["original_text_utf8_bytes"], + event_data["stored_text_utf8_bytes"], + event_data["max_tool_result_bytes"], + ) + async def add_message(self, message: dict[str, Any]) -> None: """Add a message to the context. @@ -325,6 +498,8 @@ async def add_message(self, message: dict[str, Any]) -> None: Timestamps are automatically added to message metadata for replay timing. Existing timestamps and metadata are preserved. """ + message, truncation_event = self._limit_tool_result_text_at_ingress(message) + # Add timestamp in metadata if not already present (for replay timing) existing_meta = message.get("metadata") or {} if "timestamp" not in existing_meta: @@ -347,6 +522,8 @@ async def add_message(self, message: dict[str, Any]) -> None: # Add message (no rejection - compaction happens ephemerally) self.messages.append(message) + if truncation_event is not None: + await self._emit_tool_result_ingress_truncation(truncation_event) token_count = self._estimate_tokens(self.messages) usage = token_count / self.max_tokens @@ -624,7 +801,11 @@ async def set_messages(self, messages: list[dict[str, Any]]) -> None: produced the transcript. """ restamped: list[dict[str, Any]] = [] + truncation_events: list[dict[str, Any]] = [] for i, msg in enumerate(messages): + msg, truncation_event = self._limit_tool_result_text_at_ingress(msg) + if truncation_event is not None: + truncation_events.append(truncation_event) meta = dict(msg.get("metadata") or {}) meta["_seq"] = i restamped.append({**msg, "metadata": meta}) @@ -635,6 +816,8 @@ async def set_messages(self, messages: list[dict[str, Any]]) -> None: self._stubbed_seqs = set() self._sticky_level = 0 self._last_compaction_stats = None + for truncation_event in truncation_events: + await self._emit_tool_result_ingress_truncation(truncation_event) logger.info(f"Restored {len(messages)} messages to context") async def clear(self) -> None: diff --git a/tests/test_tool_result_ingress.py b/tests/test_tool_result_ingress.py new file mode 100644 index 0000000..ba8c263 --- /dev/null +++ b/tests/test_tool_result_ingress.py @@ -0,0 +1,395 @@ +"""Regression coverage for the bounded tool-result text ingress guard.""" + +from __future__ import annotations + +import copy +import logging +from types import SimpleNamespace + +import pytest +from amplifier_core import TextBlock, ToolResult + +from amplifier_module_context_simple import SimpleContextManager, mount + + +DEFAULT_CAP = 128 * 1024 +MARKER_PREFIX = "[tool-result truncated at ingress:" +SAFE_RETRIEVAL_NOTICE = ( + "retrieve missing content using narrower read/query parameters; do not repeat " + "state-changing actions just to recover output." +) + + +def _encoded_text_bytes(text: str) -> int: + return len(text.encode("utf-8", errors="replace")) + + +class _Provider64k: + def get_model_info(self) -> SimpleNamespace: + return SimpleNamespace(context_window=64_000, max_output_tokens=2_000) + + +class _RecordingHooks: + def __init__(self) -> None: + self.events: list[tuple[str, dict]] = [] + + async def emit(self, event: str, data: dict) -> None: + self.events.append((event, dict(data))) + + +class _FailingHooks: + async def emit(self, event: str, data: dict) -> None: + raise RuntimeError("test hook failure") + + +class _Coordinator: + def __init__(self) -> None: + self.hooks = None + self.mounted: dict[str, object] = {} + + async def mount(self, kind: str, instance: object) -> None: + self.mounted[kind] = instance + + +async def _add_paired_tool_result( + context: SimpleContextManager, content: str, **extra: object +) -> None: + await context.add_message({"role": "user", "content": "inspect the local result"}) + await context.add_message( + { + "role": "assistant", + "content": "", + "tool_calls": [{"id": "call-ingress-1", "tool": "local_tool", "arguments": {}}], + } + ) + await context.add_message( + { + "role": "tool", + "name": "local_tool", + "tool_call_id": "call-ingress-1", + "content": content, + **extra, + } + ) + + +@pytest.mark.asyncio +async def test_default_caps_the_real_oversized_protected_tool_result_before_compaction(): + """The original 3.16 MB reproduction fits before request compaction runs.""" + serialized = ToolResult(success=True, output="x" * 3_163_313).get_serialized_output() + assert isinstance(serialized, str) + context = SimpleContextManager( + compaction_notice_enabled=False, + protected_tool_results=5, + ) + + await _add_paired_tool_result(context, serialized) + + stored = await context.get_messages() + stored_tool = stored[-1] + view = await context.get_messages_for_request(provider=_Provider64k()) + view_tool = next(message for message in view if message.get("role") == "tool") + + assert len(stored_tool["content"].encode("utf-8")) == DEFAULT_CAP + assert stored_tool["content"].count(MARKER_PREFIX) == 1 + assert view_tool["tool_call_id"] == "call-ingress-1" + assert view[1]["tool_calls"][0]["id"] == view_tool["tool_call_id"] + assert context._estimate_tokens(view) < 58_904 + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("output", [{"result": "x" * 140_000}, ["x" * 140_000]]) +async def test_serialized_dict_and_list_tool_outputs_are_capped_as_text(output): + serialized = ToolResult(success=True, output=output).get_serialized_output() + assert isinstance(serialized, str) + context = SimpleContextManager() + + await context.add_message({"role": "tool", "content": serialized}) + + stored = (await context.get_messages())[0]["content"] + assert len(stored.encode("utf-8")) == DEFAULT_CAP + assert stored.count(MARKER_PREFIX) == 1 + + +@pytest.mark.asyncio +async def test_string_byte_boundary_unicode_and_explicit_larger_override(): + context = SimpleContextManager() + exact = "x" * DEFAULT_CAP + one_byte_over = exact + "x" + unicode_over = "🙂" * ((DEFAULT_CAP // len("🙂".encode("utf-8"))) + 1) + + await context.add_message({"role": "tool", "content": exact}) + await context.add_message({"role": "tool", "content": one_byte_over}) + await context.add_message({"role": "tool", "content": unicode_over}) + + unchanged, clipped, unicode_clipped = await context.get_messages() + assert unchanged["content"] == exact + assert clipped["content"].count(MARKER_PREFIX) == 1 + assert len(clipped["content"].encode("utf-8")) == DEFAULT_CAP + assert unicode_clipped["content"].encode("utf-8").decode("utf-8") == unicode_clipped["content"] + assert len(unicode_clipped["content"].encode("utf-8")) <= DEFAULT_CAP + assert unicode_clipped["content"].count(MARKER_PREFIX) == 1 + + larger = SimpleContextManager(max_tool_result_bytes=DEFAULT_CAP + 1) + await larger.add_message({"role": "tool", "content": one_byte_over}) + assert (await larger.get_messages())[0]["content"] == one_byte_over + + +@pytest.mark.asyncio +async def test_lone_surrogates_do_not_crash_and_only_oversize_clipping_normalizes(): + under_cap = "prefix\ud800suffix" + over_cap = "prefix\ud800" + ("x" * DEFAULT_CAP) + context = SimpleContextManager() + + await context.add_message({"role": "tool", "content": under_cap}) + await context.add_message({"role": "tool", "content": over_cap}) + await context.add_message( + {"role": "tool", "content": [{"type": "text", "text": under_cap}]} + ) + await context.add_message( + {"role": "tool", "content": [{"type": "text", "text": over_cap}]} + ) + + string_under, string_over, block_under, block_over = await context.get_messages() + assert string_under["content"] == under_cap + assert block_under["content"][0]["text"] == under_cap + assert string_over["content"].count(MARKER_PREFIX) == 1 + assert block_over["content"][0]["text"].count(MARKER_PREFIX) == 1 + assert _encoded_text_bytes(string_over["content"]) <= DEFAULT_CAP + assert _encoded_text_bytes(block_over["content"][0]["text"]) <= DEFAULT_CAP + assert string_over["content"].encode("utf-8").decode("utf-8") == string_over["content"] + assert ( + block_over["content"][0]["text"].encode("utf-8").decode("utf-8") + == block_over["content"][0]["text"] + ) + + +@pytest.mark.parametrize("invalid", [True, 0, -1, None, "131072", 131072.0]) +def test_invalid_tool_result_ingress_caps_are_rejected(invalid): + with pytest.raises(ValueError, match="max_tool_result_bytes"): + SimpleContextManager(max_tool_result_bytes=invalid) + + +def test_cap_too_small_for_the_loud_marker_is_rejected(): + marker_at_largest_possible_length = _encoded_text_bytes( + ( + "[tool-result truncated at ingress: original_text_utf8_bytes=" + f"{__import__('sys').maxsize}; {SAFE_RETRIEVAL_NOTICE}]" + ) + ) + with pytest.raises(ValueError, match="max_tool_result_bytes"): + SimpleContextManager(max_tool_result_bytes=marker_at_largest_possible_length - 1) + + +@pytest.mark.asyncio +async def test_tool_identity_metadata_error_and_caller_objects_are_preserved(): + original = { + "role": "tool", + "name": "fetch_artifact", + "tool_call_id": "call-native-7", + "tool_use_id": "use-native-7", + "status": "error", + "error": {"message": "failed"}, + "content": "z" * (DEFAULT_CAP + 100), + "metadata": { + "openai:tool_search_items": [{"name": "glob"}], + "nested": {"keep": ["all", "of", "this"]}, + }, + } + caller_snapshot = copy.deepcopy(original) + context = SimpleContextManager() + + await context.add_message(original) + + stored = (await context.get_messages())[0] + assert original == caller_snapshot + for key in ("name", "tool_call_id", "tool_use_id", "status", "error"): + assert stored[key] == original[key] + assert stored["metadata"]["openai:tool_search_items"] == caller_snapshot["metadata"][ + "openai:tool_search_items" + ] + assert stored["metadata"]["nested"] == caller_snapshot["metadata"]["nested"] + assert MARKER_PREFIX not in str(stored["metadata"]) + + +@pytest.mark.asyncio +async def test_text_blocks_cap_aggregate_text_without_touching_media_or_unknown_blocks(): + hooks = _RecordingHooks() + image = {"type": "image", "source": {"data": "a" * (DEFAULT_CAP * 2)}} + unknown = {"type": "vendor_extension", "payload": {"keep": "unchanged"}} + original_blocks = [ + TextBlock(text="first-" + "a" * 100, visibility="user", extra_value="kept"), + image, + {"type": "text", "text": "second-" + "b" * DEFAULT_CAP, "trace": "preserve"}, + unknown, + ] + message = {"role": "tool", "content": original_blocks} + caller_snapshot = copy.deepcopy(message) + context = SimpleContextManager(hooks=hooks) + + await context.add_message(message) + + stored = (await context.get_messages())[0]["content"] + text_blocks = [ + block + for block in stored + if isinstance(block, TextBlock) or (isinstance(block, dict) and block.get("type") == "text") + ] + text = "".join( + block.text if isinstance(block, TextBlock) else block["text"] for block in text_blocks + ) + non_text = [ + block for block in stored if not (isinstance(block, TextBlock) or block.get("type") == "text") + ] + + assert message == caller_snapshot + assert _encoded_text_bytes(text) <= DEFAULT_CAP + assert text.count(MARKER_PREFIX) == 1 + assert all( + block.text if isinstance(block, TextBlock) else block["text"] for block in text_blocks + ) + assert non_text == [image, unknown] + assert isinstance(text_blocks[0], TextBlock) + assert text_blocks[0].extra_value == "kept" + assert hooks.events[0][1]["content_kind"] == "text_blocks" + + +@pytest.mark.asyncio +async def test_mixed_blocks_preserve_opaque_identity_and_drop_text_after_overflow(): + image = {"type": "image", "source": {"ref": "image-1"}} + file_ref = {"type": "file", "identity": {"id": "file-1"}} + audio_ref = {"type": "audio", "identity": {"id": "audio-1"}} + trailing_text = {"type": "text", "text": "must not survive"} + message = { + "role": "tool", + "content": [ + TextBlock(text="typed-", visibility="user"), + image, + file_ref, + audio_ref, + {"type": "text", "text": "🙂" * DEFAULT_CAP, "trace": "overflow"}, + trailing_text, + ], + } + caller_snapshot = copy.deepcopy(message) + context = SimpleContextManager() + + await context.add_message(message) + + stored = (await context.get_messages())[0]["content"] + text_blocks = [ + block + for block in stored + if isinstance(block, TextBlock) or (isinstance(block, dict) and block.get("type") == "text") + ] + stored_text = "".join( + block.text if isinstance(block, TextBlock) else block["text"] for block in text_blocks + ) + + assert message == caller_snapshot + assert stored[1] is image + assert stored[2] is file_ref + assert stored[3] is audio_ref + assert [block for block in stored if block in (image, file_ref, audio_ref)] == [ + image, + file_ref, + audio_ref, + ] + assert _encoded_text_bytes(stored_text) <= DEFAULT_CAP + assert stored_text.count(MARKER_PREFIX) == 1 + assert all( + block.text if isinstance(block, TextBlock) else block["text"] for block in text_blocks + ) + assert trailing_text not in stored + + +@pytest.mark.asyncio +async def test_image_only_and_non_tool_or_malformed_content_are_unchanged_and_silent(): + hooks = _RecordingHooks() + image_only = [{"type": "image", "source": {"base64": "x" * (DEFAULT_CAP * 2)}}] + malformed = [{"type": "text", "text": 42}, {"type": "unknown", "payload": "keep"}] + context = SimpleContextManager(hooks=hooks) + + await context.add_message({"role": "tool", "content": image_only}) + await context.add_message({"role": "tool", "content": malformed}) + await context.add_message({"role": "user", "content": "x" * (DEFAULT_CAP + 1)}) + + stored = await context.get_messages() + assert stored[0]["content"] == image_only + assert stored[1]["content"] == malformed + assert stored[2]["content"] == "x" * (DEFAULT_CAP + 1) + assert hooks.events == [] + + +@pytest.mark.asyncio +async def test_set_messages_clips_before_restamping_and_is_idempotent_on_bounded_resume(): + hooks = _RecordingHooks() + context = SimpleContextManager(hooks=hooks) + await context.add_message({"role": "tool", "content": "old" * 100_000}) + context._truncated_seqs.add(0) + resumed = [ + { + "role": "tool", + "tool_call_id": "call-resume", + "content": "x" * (DEFAULT_CAP + 2), + "metadata": {"_seq": 999, "openai:tool_search_items": [{"name": "glob"}]}, + }, + {"role": "user", "content": "continue", "metadata": {"_seq": 123}}, + ] + caller_snapshot = copy.deepcopy(resumed) + + await context.set_messages(resumed) + stored_once = await context.get_messages() + await context.set_messages(stored_once) + + stored_twice = await context.get_messages() + assert resumed == caller_snapshot + assert [message["metadata"]["_seq"] for message in stored_twice] == [0, 1] + assert context._truncated_seqs == set() + assert stored_twice[0]["metadata"]["openai:tool_search_items"] == [{"name": "glob"}] + assert stored_twice[0]["content"].count(MARKER_PREFIX) == 1 + assert len(hooks.events) == 2 # initial add plus the one oversized resumed message + + +@pytest.mark.asyncio +async def test_truncation_event_is_numeric_only_and_hook_failures_do_not_block_admission(caplog): + hooks = _RecordingHooks() + context = SimpleContextManager(hooks=hooks) + secret_text = "private-tool-output-" + "x" * DEFAULT_CAP + + await context.add_message( + { + "role": "tool", + "name": "safe_name", + "tool_call_id": "safe_call", + "content": secret_text, + } + ) + + event, data = hooks.events[0] + assert event == "context:tool_result_ingress_truncated" + assert data == { + "tool_name": "safe_name", + "tool_call_id": "safe_call", + "original_text_utf8_bytes": len(secret_text.encode("utf-8")), + "stored_text_utf8_bytes": DEFAULT_CAP, + "max_tool_result_bytes": DEFAULT_CAP, + "content_kind": "string", + } + assert secret_text not in str(data) + + failing = SimpleContextManager(hooks=_FailingHooks()) + with caplog.at_level(logging.WARNING): + await failing.add_message({"role": "tool", "content": secret_text}) + assert (await failing.get_messages())[0]["content"].count(MARKER_PREFIX) == 1 + assert SAFE_RETRIEVAL_NOTICE in caplog.text + + +@pytest.mark.asyncio +async def test_mount_forwards_the_ingress_cap(): + coordinator = _Coordinator() + + await mount(coordinator, {"max_tool_result_bytes": DEFAULT_CAP + 10}) + + assert coordinator.mounted["context"].max_tool_result_bytes == DEFAULT_CAP + 10 \ No newline at end of file