fix(cache): make proactive pruning durable and cache-aware
This commit is contained in:
+85
-13
@@ -149,6 +149,7 @@ COMPRESSED_SUMMARY_HAS_USER_TURN_KEY = "_compressed_summary_has_user_turn"
|
||||
# rolling summary, so dropping or rewriting one destroys history.
|
||||
MICRO_COMPACT_MARKER_KEY = "_micro_compact_marker"
|
||||
_DB_PERSISTED_MARKER = "_db_persisted"
|
||||
PROACTIVE_PRUNE_REARM_MODEL_CONFIG_KEY = "_proactive_prune_rearm_tokens"
|
||||
|
||||
_NO_USER_TASK_SENTINEL = "None. This session contains no user-authored turns."
|
||||
COMPRESSION_CONTINUATION_USER_CONTENT = (
|
||||
@@ -1361,6 +1362,7 @@ class ContextCompressor(ContextEngine):
|
||||
self._last_compression_telemetry = None
|
||||
self._active_compression_telemetry = None
|
||||
self._compression_telemetry_seed = None
|
||||
self._proactive_prune_rearm_tokens = 0
|
||||
|
||||
# Micro-compaction state reset
|
||||
self._micro_compact_cursor = 0
|
||||
@@ -1631,6 +1633,7 @@ class ContextCompressor(ContextEngine):
|
||||
self._last_compression_telemetry = None
|
||||
self._active_compression_telemetry = None
|
||||
self._compression_telemetry_seed = None
|
||||
self._proactive_prune_rearm_tokens = 0
|
||||
|
||||
def bind_session_state(self, session_db: Any = None, session_id: str = "") -> None:
|
||||
"""Bind the current session row so durable cooldowns can round-trip."""
|
||||
@@ -1644,9 +1647,11 @@ class ContextCompressor(ContextEngine):
|
||||
self._ineffective_compression_count = 0
|
||||
self._prellm_skip_count = 0
|
||||
self._anti_thrash_recovery_deadline = 0.0
|
||||
self._proactive_prune_rearm_tokens = 0
|
||||
self.get_active_compression_failure_cooldown()
|
||||
self._load_fallback_compression_streak()
|
||||
self._load_ineffective_compression_count()
|
||||
self._load_proactive_prune_rearm_tokens()
|
||||
|
||||
def on_session_start(self, session_id: str, **kwargs) -> None:
|
||||
"""Bind session-scoped compression state for a new or resumed session."""
|
||||
@@ -1721,6 +1726,32 @@ class ContextCompressor(ContextEngine):
|
||||
except Exception as exc:
|
||||
logger.debug("compression fallback streak lookup failed (non-sqlite): %s", exc)
|
||||
|
||||
def _load_proactive_prune_rearm_tokens(self) -> None:
|
||||
"""Restore the cache-boundary runway for a resumed durable session."""
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
session_id = getattr(self, "_session_id", "")
|
||||
getter = getattr(session_db, "get_session", None)
|
||||
if not session_id or not callable(getter):
|
||||
return
|
||||
try:
|
||||
session = getter(session_id) or {}
|
||||
raw = session.get("model_config")
|
||||
if isinstance(raw, str):
|
||||
raw = json.loads(raw) if raw.strip() else {}
|
||||
value = (
|
||||
raw.get(PROACTIVE_PRUNE_REARM_MODEL_CONFIG_KEY, 0)
|
||||
if isinstance(raw, dict)
|
||||
else 0
|
||||
)
|
||||
self._proactive_prune_rearm_tokens = max(
|
||||
0,
|
||||
int(value) if isinstance(value, (int, float, str)) else 0,
|
||||
)
|
||||
except (TypeError, ValueError, json.JSONDecodeError, sqlite3.Error) as exc:
|
||||
logger.debug("proactive prune runway lookup failed: %s", exc)
|
||||
except Exception as exc:
|
||||
logger.debug("proactive prune runway lookup failed (non-sqlite): %s", exc)
|
||||
|
||||
def _persist_fallback_compression_streak(self) -> None:
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
session_id = getattr(self, "_session_id", "")
|
||||
@@ -2080,6 +2111,7 @@ class ContextCompressor(ContextEngine):
|
||||
self._clear_compression_failure_cooldown()
|
||||
self._verify_compaction_cleared_threshold = False
|
||||
self._last_compression_made_progress = False
|
||||
self._proactive_prune_rearm_tokens = 0
|
||||
|
||||
# When the MINIMUM_CONTEXT_LENGTH floor meets/exceeds a small context
|
||||
# window, compacting at the percentage (50% → 32K of a 64K window) wastes
|
||||
@@ -2272,12 +2304,15 @@ class ContextCompressor(ContextEngine):
|
||||
# message forward. Without this gate a busy tool loop would re-fire
|
||||
# the prune nearly every iteration (each new tool pair ages an old one
|
||||
# out of the protected tail), breaking the cache per turn. Requiring a
|
||||
# meaningful batch of reclaimable tokens makes fires episodic and
|
||||
# amortized — the same way full compression is the one sanctioned
|
||||
# cache break. 0 disables the gate (commit any non-zero prune).
|
||||
# meaningful batch of reclaimable tokens, then requiring a full
|
||||
# trigger-sized growth interval before rearming, makes fires episodic
|
||||
# and amortized. 0 disables only the minimum-savings gate.
|
||||
self.proactive_prune_min_reclaim_tokens = max(
|
||||
0, int(proactive_prune_min_reclaim_tokens or 0)
|
||||
)
|
||||
# A committed prune is a prompt-cache boundary. Do not permit the next
|
||||
# one until the prompt has regrown the tokens just reclaimed.
|
||||
self._proactive_prune_rearm_tokens: int = 0
|
||||
self.min_tail_user_messages = min_tail_user_messages
|
||||
self.summary_target_ratio = max(0.10, min(summary_target_ratio, 0.80))
|
||||
self.quiet_mode = quiet_mode
|
||||
@@ -3047,12 +3082,11 @@ class ContextCompressor(ContextEngine):
|
||||
PROMPT-CACHE CONTRACT: a committed prune rewrites message bodies the
|
||||
provider has already seen, invalidating the cached prefix from the
|
||||
earliest rewritten message forward — exactly like a compression
|
||||
boundary. To keep that break episodic rather than per-turn, the prune
|
||||
only COMMITS when the estimated reclaim meets
|
||||
``proactive_prune_min_reclaim_tokens`` (measured on the actual pruned
|
||||
output, not guessed up front). Below the gate the INPUT list object is
|
||||
returned unchanged — the standard no-op caller contract (callers gate
|
||||
bookkeeping on ``result is not input``).
|
||||
boundary. A prune therefore commits only when it reclaims
|
||||
``proactive_prune_min_reclaim_tokens`` and disarms until message history
|
||||
has regrown a full trigger-sized runway. Below either gate the INPUT list
|
||||
object is returned unchanged — the standard no-op caller contract
|
||||
(callers gate bookkeeping on ``result is not input``).
|
||||
|
||||
Returns ``(messages, 0)`` — the input object — when disabled, below
|
||||
the trigger, or when the reclaim gate rejects the commit.
|
||||
@@ -3064,6 +3098,9 @@ class ContextCompressor(ContextEngine):
|
||||
# Nothing to reclaim until there are messages outside the protected tail.
|
||||
if len(messages) <= self.protect_last_n + self._protect_head_size(messages) + 1:
|
||||
return messages, 0
|
||||
before = sum(_estimate_msg_budget_tokens(m) for m in messages)
|
||||
if before < self._proactive_prune_rearm_tokens:
|
||||
return messages, 0
|
||||
pruned_msgs, pruned_count = self._prune_old_tool_results(
|
||||
messages,
|
||||
protect_tail_count=self.protect_last_n,
|
||||
@@ -3077,11 +3114,45 @@ class ContextCompressor(ContextEngine):
|
||||
# Measured-savings gate (prompt-cache hysteresis): only commit when
|
||||
# the prune reclaims a meaningful batch of tokens. Estimated on the
|
||||
# real before/after messages so dedup + arg truncation count too.
|
||||
if self.proactive_prune_min_reclaim_tokens > 0:
|
||||
before = sum(_estimate_msg_budget_tokens(m) for m in messages)
|
||||
after = sum(_estimate_msg_budget_tokens(m) for m in pruned_msgs)
|
||||
if (before - after) < self.proactive_prune_min_reclaim_tokens:
|
||||
after = sum(_estimate_msg_budget_tokens(m) for m in pruned_msgs)
|
||||
reclaimed = max(0, before - after)
|
||||
if reclaimed < self.proactive_prune_min_reclaim_tokens:
|
||||
return messages, 0
|
||||
# ``after`` includes the tool batch appended since the provider's last
|
||||
# usage reading, so both the low-water mark and future gate use the
|
||||
# same message-token estimate. Require a full trigger-sized growth
|
||||
# interval before another cache-breaking rewrite.
|
||||
runway = max(
|
||||
reclaimed,
|
||||
self.proactive_prune_tokens,
|
||||
self.proactive_prune_min_reclaim_tokens,
|
||||
)
|
||||
next_rearm_tokens = after + runway
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
session_id = getattr(self, "_session_id", "")
|
||||
if session_db and session_id:
|
||||
archive_and_compact = getattr(session_db, "archive_and_compact", None)
|
||||
if not callable(archive_and_compact):
|
||||
return messages, 0
|
||||
try:
|
||||
archive_and_compact(
|
||||
session_id,
|
||||
pruned_msgs,
|
||||
model_config_patch={
|
||||
PROACTIVE_PRUNE_REARM_MODEL_CONFIG_KEY: next_rearm_tokens,
|
||||
},
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Proactive tool-result prune DB commit failed; keeping the "
|
||||
"original transcript: %s",
|
||||
exc,
|
||||
)
|
||||
return messages, 0
|
||||
for msg in pruned_msgs:
|
||||
if isinstance(msg, dict):
|
||||
msg[_DB_PERSISTED_MARKER] = True
|
||||
self._proactive_prune_rearm_tokens = next_rearm_tokens
|
||||
return pruned_msgs, pruned_count
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -6762,6 +6833,7 @@ This compaction should PRIORITISE preserving all information related to the focu
|
||||
self._micro_compact_cursor = 0
|
||||
self._micro_compact_consecutive_failures = 0
|
||||
self._micro_compact_last_failure_cursor = -1
|
||||
self._proactive_prune_rearm_tokens = 0
|
||||
|
||||
return compressed
|
||||
|
||||
|
||||
@@ -271,6 +271,7 @@ _COMPRESSOR_ATTEMPT_STATE_FIELDS = (
|
||||
"_last_compression_telemetry",
|
||||
"_active_compression_telemetry",
|
||||
"_compression_telemetry_seed",
|
||||
"_proactive_prune_rearm_tokens",
|
||||
)
|
||||
|
||||
_COMPRESSOR_COOLDOWN_STATE_FIELDS = (
|
||||
@@ -3192,7 +3193,17 @@ def compress_context(
|
||||
# for search/recovery (Teknium review — keep one durable id
|
||||
# WITHOUT destroying history, unlike a hard replace_messages).
|
||||
# See #38763.
|
||||
agent._session_db.archive_and_compact(agent.session_id, compressed)
|
||||
from agent.context_compressor import (
|
||||
PROACTIVE_PRUNE_REARM_MODEL_CONFIG_KEY,
|
||||
)
|
||||
|
||||
agent._session_db.archive_and_compact(
|
||||
agent.session_id,
|
||||
compressed,
|
||||
model_config_patch={
|
||||
PROACTIVE_PRUNE_REARM_MODEL_CONFIG_KEY: None,
|
||||
},
|
||||
)
|
||||
split_status = "in_place_committed"
|
||||
# Reset the flush identity set so the next turn's appends are
|
||||
# diffed against the COMPACTED transcript: the compacted dicts
|
||||
@@ -3329,6 +3340,12 @@ def compress_context(
|
||||
messages[:] = copy.deepcopy(messages_before_compression)
|
||||
compressed = messages
|
||||
_compression_made_progress = False
|
||||
if "_proactive_prune_rearm_tokens" in _compressor_attempt_snapshot:
|
||||
agent.context_compressor._proactive_prune_rearm_tokens = (
|
||||
_compressor_attempt_snapshot[
|
||||
"_proactive_prune_rearm_tokens"
|
||||
]
|
||||
)
|
||||
split_status = (
|
||||
"aborted"
|
||||
if locals().get("old_session_id") is None and not in_place
|
||||
|
||||
@@ -6539,14 +6539,12 @@ def run_conversation(
|
||||
# Standard no-op caller contract: only commit when the
|
||||
# engine returned a NEW list object with a non-zero count.
|
||||
if _pruned_n and _pruned_msgs is not messages:
|
||||
# Do NOT rebuild conversation_history here. Unlike the
|
||||
# compression branch, the prune neither rotates the session
|
||||
# nor calls archive_and_compact(), so there is no new
|
||||
# persistence baseline to establish. _prune_old_tool_results
|
||||
# returns per-message copies that preserve the
|
||||
# Do NOT rebuild conversation_history here. The compressor
|
||||
# atomically rewrites the active transcript with the durable
|
||||
# rearm threshold, then stamps every returned row with
|
||||
# _DB_PERSISTED_MARKER, so the marker-based flush dedup (see
|
||||
# _flush_messages_to_session_db) already prevents both
|
||||
# duplicate writes and dropped rows. Calling
|
||||
# _flush_messages_to_session_db) prevents duplicate writes.
|
||||
# Calling
|
||||
# conversation_history_after_compression (a compaction-only
|
||||
# helper keyed on the _last_compaction_in_place flag) would be
|
||||
# a no-op at best, and on a stale in-place flag could seed
|
||||
|
||||
@@ -610,10 +610,10 @@ DEFAULT_CONFIG = {
|
||||
# itself be re-summarized.
|
||||
"proactive_prune_min_reclaim_tokens": 4096, # a proactive prune only commits
|
||||
# when it reclaims at least this many tokens
|
||||
# (measured on the pruned output). Keeps
|
||||
# prompt-cache invalidation amortized: one big
|
||||
# episodic break instead of a tiny break every
|
||||
# tool iteration. 0 = commit any non-zero prune.
|
||||
# (measured on the pruned output), then waits
|
||||
# for a full trigger-sized token runway to
|
||||
# regrow before rearming. Keeps prompt-cache
|
||||
# breaks episodic. 0 = no minimum-savings gate.
|
||||
"micro_compact": False, # opt-in: after each completed turn, fold the
|
||||
# oldest un-absorbed exchange into a rolling
|
||||
# summary, amortizing compression cost instead
|
||||
|
||||
+44
-6
@@ -6863,7 +6863,10 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
return cursor.fetchone() is not None
|
||||
|
||||
def archive_and_compact(
|
||||
self, session_id: str, compacted_messages: List[Dict[str, Any]]
|
||||
self,
|
||||
session_id: str,
|
||||
compacted_messages: List[Dict[str, Any]],
|
||||
model_config_patch: Optional[Dict[str, Any]] = None,
|
||||
) -> int:
|
||||
"""Non-destructive in-place compaction for a single durable session id.
|
||||
|
||||
@@ -6886,10 +6889,38 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
|
||||
This is the durability-preserving alternative to :meth:`replace_messages`
|
||||
for compaction. ``message_count`` is set to the ACTIVE (compacted) count,
|
||||
matching what the live load returns. Returns the new active count.
|
||||
matching what the live load returns. ``model_config_patch`` is merged
|
||||
into the session's JSON config in the same transaction; a ``None``
|
||||
value removes that key. Returns the new active count.
|
||||
"""
|
||||
|
||||
def _do(conn):
|
||||
patched_model_config = None
|
||||
if model_config_patch is not None:
|
||||
row = conn.execute(
|
||||
"SELECT model_config FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise ValueError(f"Session not found: {session_id}")
|
||||
raw = row["model_config"] if isinstance(row, sqlite3.Row) else row[0]
|
||||
config: Dict[str, Any] = {}
|
||||
if isinstance(raw, str) and raw.strip():
|
||||
try:
|
||||
parsed = json.loads(raw)
|
||||
if isinstance(parsed, dict):
|
||||
config = parsed
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
config = {}
|
||||
elif isinstance(raw, dict):
|
||||
config = dict(raw)
|
||||
for key, value in model_config_patch.items():
|
||||
if value is None:
|
||||
config.pop(key, None)
|
||||
else:
|
||||
config[key] = value
|
||||
patched_model_config = json.dumps(config) if config else None
|
||||
|
||||
# Soft-archive the live turns: active=0 hides them from the live
|
||||
# context load, compacted=1 marks them as "summarized away" (vs
|
||||
# rewind/undo's active=0+compacted=0, which means "user took it
|
||||
@@ -6906,10 +6937,17 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
)
|
||||
# message_count / tool_call_count reflect the LIVE (active) set —
|
||||
# the archived rows are still on disk but not part of the live count.
|
||||
conn.execute(
|
||||
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
|
||||
(inserted, tool_calls_total, session_id),
|
||||
)
|
||||
if model_config_patch is None:
|
||||
conn.execute(
|
||||
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
|
||||
(inserted, tool_calls_total, session_id),
|
||||
)
|
||||
else:
|
||||
conn.execute(
|
||||
"UPDATE sessions SET message_count = ?, tool_call_count = ?, "
|
||||
"model_config = ? WHERE id = ?",
|
||||
(inserted, tool_calls_total, patched_model_config, session_id),
|
||||
)
|
||||
return inserted
|
||||
|
||||
return self._execute_write(_do)
|
||||
|
||||
@@ -30,6 +30,7 @@ from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
import sqlite3
|
||||
import threading
|
||||
@@ -644,12 +645,14 @@ def test_fence_cancelled_compression_leaves_lock_reacquirable(tmp_path: Path) ->
|
||||
def _slow_summary(*_args, **_kwargs):
|
||||
summary_started.set()
|
||||
assert release_summary.wait(timeout=5)
|
||||
agent.context_compressor._proactive_prune_rearm_tokens = 0
|
||||
return [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
|
||||
agent.context_compressor.compress.side_effect = _slow_summary
|
||||
agent.context_compressor._proactive_prune_rearm_tokens = 120_000
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
fence = CompressionCommitFence()
|
||||
result = {}
|
||||
@@ -673,6 +676,7 @@ def test_fence_cancelled_compression_leaves_lock_reacquirable(tmp_path: Path) ->
|
||||
# Cancelled attempt: no mutation, and — the invariant under test — the
|
||||
# per-session compression lock is fully released.
|
||||
assert result["value"][0] is messages
|
||||
assert agent.context_compressor._proactive_prune_rearm_tokens == 120_000
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
|
||||
# The NEXT attempt (no fence — a manual /compress retry) must be able to
|
||||
@@ -840,6 +844,100 @@ def test_compression_persists_child_handoff_immediately(tmp_path: Path) -> None:
|
||||
|
||||
|
||||
|
||||
def test_rotation_publish_failure_restores_proactive_prune_runway(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
parent_sid = "PRUNE_RUNWAY_ROLLBACK_PARENT"
|
||||
db.create_session(
|
||||
parent_sid,
|
||||
source="cli",
|
||||
model_config={"keep": "value", "_proactive_prune_rearm_tokens": 120_000},
|
||||
)
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
db.append_messages_batch(parent_sid, messages)
|
||||
for message in messages:
|
||||
message["_db_persisted"] = True
|
||||
agent = _build_agent_with_db(db, parent_sid)
|
||||
agent.context_compressor._proactive_prune_rearm_tokens = 120_000
|
||||
|
||||
def _compress(*_args, **_kwargs):
|
||||
agent.context_compressor._proactive_prune_rearm_tokens = 0
|
||||
return [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
|
||||
agent.context_compressor.compress.side_effect = _compress
|
||||
durable_before = db.get_messages_as_conversation(parent_sid)
|
||||
with patch.object(
|
||||
db,
|
||||
"publish_compression_child",
|
||||
side_effect=RuntimeError("publish failed"),
|
||||
):
|
||||
returned, _sp = agent._compress_context(
|
||||
messages, "sys", approx_tokens=120_000,
|
||||
)
|
||||
|
||||
assert returned is messages
|
||||
assert agent.session_id == parent_sid
|
||||
assert agent.context_compressor._proactive_prune_rearm_tokens == 120_000
|
||||
assert db.get_messages_as_conversation(parent_sid) == durable_before
|
||||
assert json.loads(db.get_session(parent_sid)["model_config"]) == {
|
||||
"keep": "value",
|
||||
"_proactive_prune_rearm_tokens": 120_000,
|
||||
}
|
||||
|
||||
|
||||
def test_full_in_place_compression_atomically_clears_durable_prune_runway(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "IN_PLACE_CLEARS_PRUNE_RUNWAY"
|
||||
db.create_session(
|
||||
session_id,
|
||||
source="cli",
|
||||
model_config={"keep": "value", "_proactive_prune_rearm_tokens": 120_000},
|
||||
)
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
db.append_messages_batch(session_id, messages)
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
agent.compression_in_place = True
|
||||
agent.context_compressor._proactive_prune_rearm_tokens = 120_000
|
||||
|
||||
compressed, _sp = agent._compress_context(
|
||||
messages, "sys", approx_tokens=120_000,
|
||||
)
|
||||
|
||||
assert agent.session_id == session_id
|
||||
assert [message["content"] for message in db.get_messages_as_conversation(session_id)] == [
|
||||
message["content"] for message in compressed
|
||||
]
|
||||
assert json.loads(db.get_session(session_id)["model_config"]) == {"keep": "value"}
|
||||
|
||||
|
||||
def test_rotation_child_starts_without_durable_prune_runway(tmp_path: Path) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
parent_sid = "ROTATION_CLEARS_PRUNE_RUNWAY"
|
||||
db.create_session(
|
||||
parent_sid,
|
||||
source="cli",
|
||||
model_config={"keep": "parent", "_proactive_prune_rearm_tokens": 120_000},
|
||||
)
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
db.append_messages_batch(parent_sid, messages)
|
||||
agent = _build_agent_with_db(db, parent_sid)
|
||||
|
||||
agent._compress_context(messages, "sys", approx_tokens=120_000)
|
||||
|
||||
assert agent.session_id != parent_sid
|
||||
child_config = json.loads(db.get_session(agent.session_id)["model_config"])
|
||||
assert "_proactive_prune_rearm_tokens" not in child_config
|
||||
assert json.loads(db.get_session(parent_sid)["model_config"])[
|
||||
"_proactive_prune_rearm_tokens"
|
||||
] == 120_000
|
||||
|
||||
|
||||
@pytest.mark.parametrize("in_place", [False, True])
|
||||
def test_equal_copy_compression_result_does_not_rewrite_session(
|
||||
tmp_path: Path,
|
||||
|
||||
@@ -0,0 +1,211 @@
|
||||
"""Restart-safety regressions for proactive tool-result pruning."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.context_compressor import _estimate_msg_budget_tokens
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
_REARM_KEY = "_proactive_prune_rearm_tokens"
|
||||
|
||||
|
||||
def _assistant_call(call_id: str) -> dict:
|
||||
return {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [{
|
||||
"id": call_id,
|
||||
"type": "function",
|
||||
"function": {"name": "terminal", "arguments": '{"cmd":"ls"}'},
|
||||
}],
|
||||
}
|
||||
|
||||
|
||||
def _tool_result(call_id: str, content: str) -> dict:
|
||||
return {"role": "tool", "tool_call_id": call_id, "content": content}
|
||||
|
||||
|
||||
def _history(*, large_chars: int = 24_000) -> list[dict]:
|
||||
messages: list[dict] = [{"role": "user", "content": "start"}]
|
||||
for index in range(8):
|
||||
call_id = f"call_{index}"
|
||||
messages.append(_assistant_call(call_id))
|
||||
content = chr(65 + index) * large_chars if index < 3 else "ok"
|
||||
messages.append(_tool_result(call_id, content))
|
||||
return messages
|
||||
|
||||
|
||||
def _build_agent(db: SessionDB, session_id: str, *, platform: str = "telegram"):
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
|
||||
return AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=db,
|
||||
session_id=session_id,
|
||||
platform=platform,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
|
||||
def _configure_pruning(agent) -> None:
|
||||
compressor = agent.context_compressor
|
||||
compressor.proactive_prune_tokens = 48_000
|
||||
compressor.proactive_prune_min_result_chars = 8_000
|
||||
compressor.proactive_prune_min_reclaim_tokens = 4_096
|
||||
compressor.protect_first_n = 2
|
||||
compressor.protect_last_n = 4
|
||||
|
||||
|
||||
def _model_config(db: SessionDB, session_id: str) -> dict:
|
||||
raw = db.get_session(session_id)["model_config"]
|
||||
return json.loads(raw) if raw else {}
|
||||
|
||||
|
||||
def test_gateway_eviction_reload_keeps_prune_and_durable_runway(tmp_path: Path) -> None:
|
||||
"""A fresh gateway agent must reload both the pruned body and its runway."""
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "GATEWAY_PRUNE_RESTART"
|
||||
db.create_session(
|
||||
session_id, source="telegram", model_config={"keep": "value"},
|
||||
)
|
||||
db.append_messages_batch(session_id, _history())
|
||||
|
||||
first_agent = _build_agent(db, session_id)
|
||||
_configure_pruning(first_agent)
|
||||
before = db.get_messages_as_conversation(session_id)
|
||||
pruned, count = first_agent.context_compressor.prune_tool_results_only(
|
||||
before, current_tokens=120_000,
|
||||
)
|
||||
|
||||
assert count >= 1
|
||||
durable = db.get_messages_as_conversation(session_id)
|
||||
assert [message["content"] for message in durable] == [
|
||||
message["content"] for message in pruned
|
||||
]
|
||||
assert len(durable[2]["content"]) < 24_000
|
||||
stored_runway = _model_config(db, session_id)[_REARM_KEY]
|
||||
assert _model_config(db, session_id)["keep"] == "value"
|
||||
assert stored_runway > sum(map(_estimate_msg_budget_tokens, durable))
|
||||
|
||||
# Simulate gateway cache eviction / process restart: construct a wholly
|
||||
# new AIAgent and load the active transcript from SQLite.
|
||||
resumed_agent = _build_agent(db, session_id)
|
||||
_configure_pruning(resumed_agent)
|
||||
assert resumed_agent.context_compressor._proactive_prune_rearm_tokens == stored_runway
|
||||
reloaded = db.get_messages_as_conversation(session_id)
|
||||
archived_before = len(db.get_messages(session_id, include_inactive=True))
|
||||
result, second_count = resumed_agent.context_compressor.prune_tool_results_only(
|
||||
reloaded, current_tokens=1_000_000,
|
||||
)
|
||||
|
||||
assert result is reloaded
|
||||
assert second_count == 0
|
||||
assert len(db.get_messages(session_id, include_inactive=True)) == archived_before
|
||||
|
||||
|
||||
def test_fresh_agent_rearms_after_durable_history_regrowth_once(tmp_path: Path) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "PRUNE_DURABLE_REGROWTH"
|
||||
db.create_session(session_id, source="telegram")
|
||||
db.append_messages_batch(session_id, _history())
|
||||
first_agent = _build_agent(db, session_id)
|
||||
_configure_pruning(first_agent)
|
||||
first, first_count = first_agent.context_compressor.prune_tool_results_only(
|
||||
db.get_messages_as_conversation(session_id), current_tokens=120_000,
|
||||
)
|
||||
assert first_count >= 1
|
||||
first_runway = _model_config(db, session_id)[_REARM_KEY]
|
||||
|
||||
growth = [
|
||||
_assistant_call("regrown_large"),
|
||||
_tool_result("regrown_large", "z" * 240_000),
|
||||
_assistant_call("tail_1"),
|
||||
_tool_result("tail_1", "ok"),
|
||||
_assistant_call("tail_2"),
|
||||
_tool_result("tail_2", "ok"),
|
||||
]
|
||||
db.append_messages_batch(session_id, growth)
|
||||
|
||||
resumed = _build_agent(db, session_id)
|
||||
_configure_pruning(resumed)
|
||||
grown = db.get_messages_as_conversation(session_id)
|
||||
assert sum(map(_estimate_msg_budget_tokens, grown)) >= first_runway
|
||||
second, second_count = resumed.context_compressor.prune_tool_results_only(
|
||||
grown, current_tokens=1_000_000,
|
||||
)
|
||||
|
||||
assert second_count >= 1
|
||||
second_runway = _model_config(db, session_id)[_REARM_KEY]
|
||||
assert second_runway > first_runway
|
||||
|
||||
restarted = _build_agent(db, session_id)
|
||||
_configure_pruning(restarted)
|
||||
durable = db.get_messages_as_conversation(session_id)
|
||||
result, third_count = restarted.context_compressor.prune_tool_results_only(
|
||||
durable, current_tokens=1_000_000,
|
||||
)
|
||||
assert result is durable
|
||||
assert third_count == 0
|
||||
assert restarted.context_compressor._proactive_prune_rearm_tokens == second_runway
|
||||
|
||||
|
||||
def test_prune_persistence_failure_is_a_noop(tmp_path: Path) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "PRUNE_PERSISTENCE_FAILURE"
|
||||
db.create_session(session_id, source="telegram")
|
||||
db.append_messages_batch(session_id, _history())
|
||||
agent = _build_agent(db, session_id)
|
||||
_configure_pruning(agent)
|
||||
messages = db.get_messages_as_conversation(session_id)
|
||||
original_contents = [message["content"] for message in messages]
|
||||
|
||||
with patch.object(
|
||||
db, "archive_and_compact", side_effect=RuntimeError("disk full"),
|
||||
):
|
||||
result, count = agent.context_compressor.prune_tool_results_only(
|
||||
messages, current_tokens=120_000,
|
||||
)
|
||||
|
||||
assert result is messages
|
||||
assert count == 0
|
||||
assert agent.context_compressor._proactive_prune_rearm_tokens == 0
|
||||
assert [message["content"] for message in messages] == original_contents
|
||||
assert [message["content"] for message in db.get_messages_as_conversation(session_id)] == original_contents
|
||||
assert _REARM_KEY not in _model_config(db, session_id)
|
||||
|
||||
|
||||
def test_archive_model_config_patch_rolls_back_with_transcript(tmp_path: Path) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "PRUNE_ATOMIC_ARCHIVE_FAILURE"
|
||||
db.create_session(
|
||||
session_id,
|
||||
source="telegram",
|
||||
model_config={"keep": "value", _REARM_KEY: 120_000},
|
||||
)
|
||||
original = [{"role": "user", "content": "original"}]
|
||||
db.append_messages_batch(session_id, original)
|
||||
|
||||
with patch.object(
|
||||
db, "_insert_message_rows", side_effect=RuntimeError("insert failed"),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="insert failed"):
|
||||
db.archive_and_compact(
|
||||
session_id,
|
||||
[{"role": "user", "content": "replacement"}],
|
||||
model_config_patch={_REARM_KEY: None},
|
||||
)
|
||||
|
||||
assert db.get_messages_as_conversation(session_id)[0]["content"] == "original"
|
||||
assert _model_config(db, session_id) == {"keep": "value", _REARM_KEY: 120_000}
|
||||
@@ -11,7 +11,11 @@ Mirrors the construction/patching conventions in test_context_compressor.py.
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
from agent.context_compressor import ContextCompressor, _PRUNED_TOOL_PLACEHOLDER
|
||||
from agent.context_compressor import (
|
||||
ContextCompressor,
|
||||
_PRUNED_TOOL_PLACEHOLDER,
|
||||
_estimate_msg_budget_tokens,
|
||||
)
|
||||
|
||||
LARGE_WINDOW = 1_000_000
|
||||
|
||||
@@ -98,11 +102,78 @@ def test_idempotent():
|
||||
msgs = _build(8, big_indices={0, 1, 2})
|
||||
first, n1 = c.prune_tool_results_only(msgs, current_tokens=120_000)
|
||||
assert n1 >= 3
|
||||
second, n2 = c.prune_tool_results_only(first, current_tokens=120_000)
|
||||
# No usage reading bypasses the token gate and exercises prune idempotence.
|
||||
second, n2 = c.prune_tool_results_only(first, current_tokens=None)
|
||||
assert n2 == 0
|
||||
assert [m.get("content") for m in second] == [m.get("content") for m in first]
|
||||
|
||||
|
||||
def test_rearms_only_after_reclaimed_token_runway():
|
||||
"""A prune boundary must earn back its cache break before the next one."""
|
||||
c = _compressor(
|
||||
proactive_prune_tokens=48_000,
|
||||
proactive_prune_min_result_chars=8_000,
|
||||
)
|
||||
msgs = _build(8, big_indices={0, 1, 2, 6, 7})
|
||||
|
||||
first, n1 = c.prune_tool_results_only(msgs, current_tokens=120_000)
|
||||
assert n1 >= 3
|
||||
rearm_tokens = sum(map(_estimate_msg_budget_tokens, first)) + 48_000
|
||||
|
||||
# Age the two protected large results out of the tail. They are now a
|
||||
# valid >=4K-token prune candidate, but the post-prune prompt has not yet
|
||||
# regrown the tokens reclaimed at the first cache-breaking boundary.
|
||||
grown = first + [
|
||||
_assistant_call("call_8"),
|
||||
_tool_msg("call_8", "ok"),
|
||||
_assistant_call("call_9"),
|
||||
_tool_msg("call_9", "ok"),
|
||||
]
|
||||
assert sum(map(_estimate_msg_budget_tokens, grown)) < rearm_tokens
|
||||
blocked, n2 = c.prune_tool_results_only(grown, current_tokens=1_000_000)
|
||||
assert n2 == 0
|
||||
assert blocked is grown
|
||||
assert len(_tool_by_id(blocked, "call_6")["content"]) == 9000
|
||||
assert len(_tool_by_id(blocked, "call_7")["content"]) == 9000
|
||||
|
||||
missing = rearm_tokens - sum(map(_estimate_msg_budget_tokens, grown))
|
||||
regrown = grown + [{"role": "user", "content": "x" * (missing * 4)}]
|
||||
assert sum(map(_estimate_msg_budget_tokens, regrown)) >= rearm_tokens
|
||||
rearmed, n3 = c.prune_tool_results_only(regrown, current_tokens=1_000_000)
|
||||
assert n3 >= 2
|
||||
assert rearmed is not regrown
|
||||
|
||||
|
||||
def test_successful_full_compression_resets_proactive_runway():
|
||||
"""A full compression establishes a fresh cache boundary and baseline."""
|
||||
c = _compressor(
|
||||
proactive_prune_tokens=48_000,
|
||||
proactive_prune_min_result_chars=8_000,
|
||||
)
|
||||
first, n1 = c.prune_tool_results_only(
|
||||
_build(8, big_indices={0, 1, 2}), current_tokens=120_000,
|
||||
)
|
||||
assert n1 >= 3
|
||||
|
||||
history = [{"role": "system", "content": "sys"}]
|
||||
for i in range(12):
|
||||
history.append({
|
||||
"role": "user" if i % 2 == 0 else "assistant",
|
||||
"content": f"turn {i} " + ("x" * 1000),
|
||||
})
|
||||
c.tail_token_budget = 50
|
||||
with patch.object(c, "_generate_summary", return_value="summary"):
|
||||
compressed = c.compress(history, current_tokens=500_000, force=True)
|
||||
assert c._last_compression_made_progress is True
|
||||
assert len(compressed) < len(history)
|
||||
|
||||
# The successful full boundary supersedes the old proactive-prune runway.
|
||||
fresh = _build(8, big_indices={0, 1, 2})
|
||||
result, pruned = c.prune_tool_results_only(fresh, current_tokens=48_000)
|
||||
assert pruned >= 3
|
||||
assert result is not fresh
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -136,6 +136,25 @@ def _run_tool_loop(agent, n_tool_iterations: int):
|
||||
|
||||
|
||||
class TestProactivePruneLoopWiring:
|
||||
def test_full_compression_preempts_proactive_prune(self, agent):
|
||||
agent.context_compressor.should_compress.return_value = True
|
||||
|
||||
def _compress(messages, system_message, **_kwargs):
|
||||
return [dict(m) for m in messages], system_message
|
||||
|
||||
with (
|
||||
patch.object(agent, "_compress_context", side_effect=_compress) as compress,
|
||||
patch(
|
||||
"agent.conversation_loop.conversation_history_after_compression",
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
result = _run_tool_loop(agent, n_tool_iterations=1)
|
||||
|
||||
assert result["completed"] is True
|
||||
compress.assert_called_once()
|
||||
agent.context_compressor.prune_tool_results_only.assert_not_called()
|
||||
|
||||
def test_prune_consulted_when_compression_stands_down(self, agent):
|
||||
calls = []
|
||||
|
||||
|
||||
Reference in New Issue
Block a user