refactor(agent): finish memory/compaction/prompt-cache compaction pass (>=25% LOC)

This commit is contained in:
Teknium
2026-09-02 19:50:13 -07:00
parent c629274efc
commit 3b5aa80473
6 changed files with 150 additions and 299 deletions
+38 -83
View File
@@ -56,11 +56,9 @@ def _accepts_require_checkpoint(fn: Callable[..., Any]) -> bool:
params = _signature_params(fn)
if params is None:
return False
if _has_var_kwargs(params):
return True
param = params.get("require_checkpoint")
return param is not None and param.kind in (
inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD,
return _has_var_kwargs(params) or (
param is not None and param.kind in (inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD)
)
@@ -91,9 +89,7 @@ def normalize_tool_schema(schema: Any) -> Optional[Dict[str, Any]]:
def memory_provider_tools_enabled(
enabled_toolsets: Optional[List[str]],
disabled_toolsets: Optional[List[str]] = None,
*,
enabled_toolsets: Optional[List[str]], disabled_toolsets: Optional[List[str]] = None, *,
memory_tool_present: bool = False,
) -> bool:
"""Return whether external memory-provider tools should be exposed."""
@@ -139,7 +135,6 @@ def inject_memory_provider_tools(agent: Any) -> int:
if not memory_manager or tools is None:
return 0
existing_tool_names = {_tool_name(tool) for tool in tools if isinstance(tool, dict)}
if not memory_provider_tools_exposed(agent):
# Say so once: a silent 0 leaves the provider looking "half on" with no clue which
# config key (platform_toolsets / disabled_toolsets) gated it.
@@ -158,10 +153,9 @@ def inject_memory_provider_tools(agent: Any) -> int:
if not callable(get_schemas):
return 0
valid_tool_names = getattr(agent, "valid_tool_names", None)
if valid_tool_names is None:
valid_tool_names = agent.valid_tool_names = set()
if getattr(agent, "valid_tool_names", None) is None:
agent.valid_tool_names = set()
existing_tool_names = {_tool_name(tool) for tool in tools if isinstance(tool, dict)}
added = 0
for raw_schema in get_schemas():
schema = normalize_tool_schema(raw_schema)
@@ -170,24 +164,18 @@ def inject_memory_provider_tools(agent: Any) -> int:
"Memory provider returned a tool schema with no resolvable "
"name; skipping to avoid poisoning the request (%r)", raw_schema,
)
continue
tool_name = schema["name"]
if tool_name in existing_tool_names:
continue
tools.append({"type": "function", "function": schema})
valid_tool_names.add(tool_name)
existing_tool_names.add(tool_name)
added += 1
elif schema["name"] not in existing_tool_names:
tools.append({"type": "function", "function": schema})
agent.valid_tool_names.add(schema["name"])
existing_tool_names.add(schema["name"])
added += 1
return added
# -- Context fencing helpers --------------------------------------------------
_FENCE_TAG_RE = re.compile(r'</?\s*memory-context\s*>', re.IGNORECASE)
_INTERNAL_CONTEXT_RE = re.compile(
r'<\s*memory-context\s*>[\s\S]*?</\s*memory-context\s*>',
re.IGNORECASE,
)
_INTERNAL_CONTEXT_RE = re.compile(r'<\s*memory-context\s*>[\s\S]*?</\s*memory-context\s*>', re.IGNORECASE)
_INTERNAL_NOTE_RE = re.compile(
r'\[System note:\s*The following is recalled memory context,\s*NOT new user input\.\s*Treat as (?:informational background data|authoritative reference data[^\]]*)\.\]\s*',
re.IGNORECASE,
@@ -196,9 +184,9 @@ _INTERNAL_NOTE_RE = re.compile(
def sanitize_context(text: str) -> str:
"""Strip fence tags, injected context blocks, and system notes from provider output."""
text = _INTERNAL_CONTEXT_RE.sub('', text)
text = _INTERNAL_NOTE_RE.sub('', text)
return _FENCE_TAG_RE.sub('', text)
for pattern in (_INTERNAL_CONTEXT_RE, _INTERNAL_NOTE_RE, _FENCE_TAG_RE):
text = pattern.sub('', text)
return text
class StreamingContextScrubber:
@@ -228,27 +216,25 @@ class StreamingContextScrubber:
buf = self._buf + text
self._buf = ""
out: list[str] = []
while buf:
if self._in_span:
idx = buf.lower().find(self._CLOSE_TAG)
if idx == -1:
# Hold back a potential partial close tag; drop the rest.
held = self._max_partial_suffix(buf, self._CLOSE_TAG)
self._buf = buf[-held:] if held else ""
break
buf = buf[idx + len(self._CLOSE_TAG):]
held = self._max_partial_suffix(buf, self._CLOSE_TAG) # potential partial close tag
tag = self._CLOSE_TAG
else:
idx = self._find_boundary_open_tag(buf)
if idx == -1:
held = self._max_pending_open_suffix(buf) or self._max_partial_suffix(buf, self._OPEN_TAG)
held = self._max_pending_open_suffix(buf) or self._max_partial_suffix(buf, self._OPEN_TAG)
tag = self._OPEN_TAG
if idx == -1:
# Hold back the possible partial tag; inside a span the rest is dropped.
if not self._in_span:
self._append_visible(out, buf[:-held] if held else buf)
self._buf = buf[-held:] if held else ""
break
self._buf = buf[-held:] if held else ""
break
if not self._in_span:
self._append_visible(out, buf[:idx])
buf = buf[idx + len(self._OPEN_TAG):]
buf = buf[idx + len(tag):]
self._in_span = not self._in_span
return "".join(out)
def flush(self) -> str:
@@ -356,8 +342,6 @@ class MemoryManager:
"status": "not_started", "abandoned_writes": 0, "abandoned_prefetches": 0, "active_tasks": 0,
}
# -- Fan-out helper ------------------------------------------------------
def _each_provider(
self, label: str, call: Callable[[MemoryProvider], Any], *,
level: int = logging.DEBUG, providers: Optional[List[MemoryProvider]] = None, exc_info: bool = False,
@@ -375,8 +359,6 @@ class MemoryManager:
logger.log(level, "Memory provider '%s' %s: %s", provider.name, label, e, exc_info=exc_info)
return results
# -- Registration --------------------------------------------------------
def add_provider(self, provider: MemoryProvider) -> None:
"""Register a provider; builtin always accepted, only ONE external allowed."""
if provider.name != "builtin":
@@ -425,8 +407,6 @@ class MemoryManager:
def get_provider(self, name: str) -> Optional[MemoryProvider]:
return next((p for p in self._providers if p.name == name), None)
# -- System prompt -------------------------------------------------------
def build_system_prompt(self) -> str:
"""Join every provider's non-empty ``system_prompt_block()`` with blank lines."""
blocks = self._each_provider(
@@ -434,8 +414,6 @@ class MemoryManager:
)
return "\n\n".join(b for b in blocks if b)
# -- Prefetch / recall ---------------------------------------------------
# A /skill or /bundle turn embeds the whole skill body in the model-facing message;
# providers get just the user's instruction (None for a bare invocation).
_strip_skill_scaffolding = staticmethod(extract_user_instruction_from_skill_message)
@@ -528,8 +506,6 @@ class MemoryManager:
kind="prefetch",
)
# -- Sync ----------------------------------------------------------------
@staticmethod
def _provider_sync_accepts_messages(provider: MemoryProvider) -> bool:
"""Whether ``sync_turn`` accepts a ``messages`` keyword (uninspectable → assume yes)."""
@@ -561,8 +537,6 @@ class MemoryManager:
lambda: self._each_provider("sync_turn failed", _sync, level=logging.WARNING, providers=providers)
)
# -- Background dispatch -------------------------------------------------
def _submit_background(self, fn, *, kind: str = "write") -> None:
"""Queue ``fn`` on the serialized worker and track its durability class.
@@ -588,11 +562,11 @@ class MemoryManager:
return
if future is not None:
future.add_done_callback(self._forget_background_future)
return
try:
fn()
except Exception as e: # pragma: no cover - fn guards internally
logger.debug("Inline memory background task failed: %s", e)
else:
try:
fn()
except Exception as e: # pragma: no cover - fn guards internally
logger.debug("Inline memory background task failed: %s", e)
def _forget_background_future(self, future: Future) -> None:
with self._sync_executor_lock:
@@ -605,9 +579,7 @@ class MemoryManager:
if self._sync_executor is not None:
return self._sync_executor
with self._sync_executor_lock:
if self._shutting_down:
return None
if self._sync_executor is None:
if self._sync_executor is None and not self._shutting_down:
try:
# Daemon workers: a wedged provider must never block interpreter exit.
from tools.daemon_pool import DaemonThreadPoolExecutor
@@ -625,17 +597,13 @@ class MemoryManager:
if executor is None:
return True
try:
fut = executor.submit(lambda: None)
executor.submit(lambda: None).result(timeout=timeout)
except RuntimeError:
return True # executor already shut down — nothing pending
try:
fut.result(timeout=timeout)
except Exception:
return False
return True
# -- Tools ---------------------------------------------------------------
def get_all_tool_schemas(self) -> List[Dict[str, Any]]:
"""Collect deduplicated tool schemas from all providers.
@@ -654,11 +622,9 @@ class MemoryManager:
"Memory provider '%s' returned a tool schema with "
"no resolvable name; skipping (%r)", provider.name, raw_schema,
)
continue
name = schema["name"]
if name not in _HERMES_CORE_TOOLS and name not in seen:
elif schema["name"] not in _HERMES_CORE_TOOLS and schema["name"] not in seen:
schemas.append(schema)
seen.add(name)
seen.add(schema["name"])
self._each_provider("get_tool_schemas() failed", _collect, level=logging.WARNING)
return schemas
@@ -680,14 +646,10 @@ class MemoryManager:
logger.error("Memory provider '%s' handle_tool_call(%s) failed: %s", provider.name, tool_name, e)
return tool_error(f"Memory tool '{tool_name}' failed: {e}")
# -- Lifecycle hooks -----------------------------------------------------
def on_turn_start(self, turn_number: int, message: str, **kwargs) -> None:
"""Notify all providers of a new turn (kwargs: remaining_tokens, model, platform, tool_count)."""
self._each_provider("on_turn_start failed", lambda p: p.on_turn_start(turn_number, message, **kwargs))
def on_session_end(self, messages: List[Dict[str, Any]]) -> None:
"""Notify all providers of session end."""
self._each_provider(
"on_session_end failed", lambda p: p.on_session_end(messages), level=logging.WARNING, exc_info=True,
)
@@ -839,9 +801,7 @@ class MemoryManager:
result = json.loads(result)
except Exception:
return False
if not isinstance(result, dict):
return False
return result.get("success") is True and result.get("staged") is not True
return isinstance(result, dict) and result.get("success") is True and result.get("staged") is not True
def notify_memory_tool_write(
self, tool_result: Any, tool_args: Dict[str, Any], *,
@@ -855,16 +815,12 @@ class MemoryManager:
"""
if not self._memory_tool_result_succeeded(tool_result):
return
target = str(tool_args.get("target") or "memory")
operations = tool_args.get("operations")
if not (isinstance(operations, list) and operations):
operations = [{k: tool_args.get(k) for k in ("action", "content", "old_text")}]
operations = [tool_args]
for op in operations:
if not isinstance(op, dict):
continue
action = str(op.get("action") or "")
action = str(op.get("action") or "") if isinstance(op, dict) else ""
if action not in self._MIRRORED_MEMORY_ACTIONS:
continue
try:
@@ -877,7 +833,6 @@ class MemoryManager:
logger.debug("notify_memory_tool_write failed for op %s: %s", action, e)
def on_delegation(self, task: str, result: str, *, child_session_id: str = "", **kwargs) -> None:
"""Notify all providers that a subagent completed."""
self._each_provider(
"on_delegation failed",
lambda p: p.on_delegation(task, result, child_session_id=child_session_id, **kwargs),
+22 -54
View File
@@ -71,8 +71,7 @@ class MemoryProvider(ABC):
@abstractmethod
def is_available(self) -> bool:
"""Configured, credentialed and ready? Gates activation at agent init;
check config/deps only — no network calls."""
"""Configured, credentialed and ready? Gates activation; check config/deps only, no network."""
@abstractmethod
def initialize(self, session_id: str, **kwargs) -> None:
@@ -85,13 +84,11 @@ class MemoryProvider(ABC):
"""
def unavailable_reason(self) -> str:
"""Short user-facing hint for the "provider unavailable" warning (e.g.
which package to install); ``initialize()`` never runs when unavailable."""
"""User-facing hint for the "provider unavailable" warning (``initialize()`` never runs then)."""
return ""
def system_prompt_block(self) -> str:
"""STATIC system-prompt text (instructions, status); "" to skip.
Recalled context goes through prefetch(), not here."""
"""STATIC system-prompt text; "" to skip. Recalled context goes through prefetch(), not here."""
return ""
def prefetch(self, query: str, *, session_id: str = "") -> str:
@@ -103,26 +100,19 @@ class MemoryProvider(ABC):
"""Queue a background recall after each turn; prefetch() consumes it next turn."""
def recall_status(self) -> Optional[RecallStatus]:
"""What the most recent :meth:`prefetch` injected, for a deterministic
"recalled N memories" indicator. ``None`` = nothing / no indicator.
Must reflect only the LAST prefetch, never a stale prior count."""
"""What the most recent :meth:`prefetch` injected (``None`` = no indicator). Must reflect
only the LAST prefetch, never a stale prior count."""
return None
def sync_turn(
self,
user_content: str,
assistant_content: str,
*,
session_id: str = "",
messages: Optional[List[Dict[str, Any]]] = None,
self, user_content: str, assistant_content: str, *,
session_id: str = "", messages: Optional[List[Dict[str, Any]]] = None,
) -> None:
"""Persist a completed turn; should be non-blocking. ``messages`` is the
OpenAI-style list as of this turn, including tool calls/results."""
"""Persist a completed turn (non-blocking). ``messages`` is the OpenAI-style list so far."""
@abstractmethod
def get_tool_schemas(self) -> List[Dict[str, Any]]:
"""OpenAI function-calling schemas ({"name", "description", "parameters"});
[] for context-only providers."""
"""OpenAI function-calling schemas ({"name", "description", "parameters"}); [] if none."""
def handle_tool_call(self, tool_name: str, args: Dict[str, Any], **kwargs) -> str:
"""Handle one of this provider's tools; must return a JSON string."""
@@ -134,21 +124,13 @@ class MemoryProvider(ABC):
# -- Optional hooks (override to opt in) ---------------------------------
def on_turn_start(self, turn_number: int, message: str, **kwargs) -> None:
"""Per-turn tick (turn-counting, scope management, maintenance).
kwargs may include remaining_tokens, model, platform, tool_count."""
"""Per-turn tick. kwargs may include remaining_tokens, model, platform, tool_count."""
def on_session_end(self, messages: List[Dict[str, Any]]) -> None:
"""End-of-session extraction over the full history. Fires only at real
session boundaries (CLI exit, /reset, gateway expiry), never per-turn."""
"""End-of-session extraction; fires only at real session boundaries, never per-turn."""
def on_session_switch(
self,
new_session_id: str,
*,
parent_session_id: str = "",
reset: bool = False,
rewound: bool = False,
**kwargs,
self, new_session_id: str, *, parent_session_id: str = "", reset: bool = False, rewound: bool = False, **kwargs,
) -> None:
"""session_id reassigned mid-process (/resume, /branch, /reset, /new, compression)
without teardown: rebind per-session state so later writes land in the right record.
@@ -156,14 +138,11 @@ class MemoryProvider(ABC):
same id but the transcript was truncated."""
def on_pre_compress(self, messages: List[Dict[str, Any]]) -> str:
"""Extract insights from ``messages`` about to be compressed; the returned
text is fed into the compression summary prompt ("" = nothing)."""
"""Extract insights from ``messages`` about to be compressed, fed into the summary prompt."""
return ""
def on_delegation(self, task: str, result: str, *,
child_session_id: str = "", **kwargs) -> None:
"""PARENT-side observation of a completed delegation (task prompt + final
result); the subagent itself has no provider session (skip_memory=True)."""
def on_delegation(self, task: str, result: str, *, child_session_id: str = "", **kwargs) -> None:
"""PARENT-side observation of a completed delegation (the subagent has no provider session)."""
def get_config_schema(self) -> List[Dict[str, Any]]:
"""Setup fields for ``hermes memory setup`` ([] if none): ``key``, ``description``,
@@ -173,25 +152,14 @@ class MemoryProvider(ABC):
return []
def save_config(self, values: Dict[str, Any], hermes_home: str) -> None:
"""Write non-secret setup ``values`` (secrets go to .env) to the provider's
native config location. Plugins MUST either override this or use only
env vars (every schema field carrying ``env_var``) and keep the no-op."""
"""Write non-secret setup ``values`` to the provider's native config. Plugins MUST either
override this or use only env vars (every schema field carrying ``env_var``)."""
def on_memory_write(
self,
action: str,
target: str,
content: str,
metadata: Optional[Dict[str, Any]] = None,
) -> None:
"""Mirror a built-in memory-tool write. ``action`` is add | replace |
remove, ``target`` is memory | user; ``metadata`` (when available) has
provenance such as write_origin, execution_context, session_id,
parent_session_id, platform, tool_name."""
def on_memory_write(self, action: str, target: str, content: str, metadata: Optional[Dict[str, Any]] = None) -> None:
"""Mirror a built-in memory-tool write (``action``: add | replace | remove; ``target``:
memory | user; ``metadata``: provenance such as write_origin, session_id, tool_name)."""
def backup_paths(self) -> List[str]:
"""Absolute paths of provider state OUTSIDE HERMES_HOME (e.g. ``~/.honcho``)
so ``hermes backup``/``hermes import`` can capture and restore them; paths
outside the home dir are skipped. MUST work without ``initialize()`` or
network — resolve from config/env."""
"""Absolute paths of provider state OUTSIDE HERMES_HOME for ``hermes backup``/``import``
(paths outside the home dir are skipped). MUST work without ``initialize()`` or network."""
return []
+25 -47
View File
@@ -11,6 +11,7 @@ import hashlib
import json
import logging
import re
from functools import partial
from typing import Any, Callable
logger = logging.getLogger(__name__)
@@ -49,22 +50,15 @@ def _fix_str_field(container: Any, key: Any, fix: Callable[[str], str]) -> bool:
def _sanitize_structure(payload: Any, fix: Callable[[str], str]) -> bool:
"""Apply ``fix`` to every str inside nested dict/list ``payload`` in-place."""
found = False
def _walk(node):
nonlocal found
if isinstance(node, dict):
items = list(node.items())
elif isinstance(node, list):
items = list(enumerate(node))
else:
return
for key, value in items:
stack = [payload]
while stack:
node = stack.pop()
items = node.items() if isinstance(node, dict) else enumerate(node) if isinstance(node, list) else ()
for key, value in list(items):
if isinstance(value, str):
found |= _fix_str_field(node, key, fix)
elif isinstance(value, (dict, list)):
_walk(value)
_walk(payload)
stack.append(value)
return found
@@ -106,29 +100,13 @@ def _sanitize_messages(messages: list, fix: Callable[[str], str], *, deep: bool)
return found
def _sanitize_structure_surrogates(payload: Any) -> bool:
"""Replace surrogates in nested dict/list payloads in-place; True if any replaced."""
return _sanitize_structure(payload, _sanitize_surrogates)
def _sanitize_messages_surrogates(messages: list) -> bool:
"""Replace surrogates in all string content of a messages list in-place; True if any found."""
return _sanitize_messages(messages, _sanitize_surrogates, deep=True)
def _sanitize_structure_non_ascii(payload: Any) -> bool:
"""Strip non-ASCII from nested dict/list payloads in-place; True if any stripped."""
return _sanitize_structure(payload, _strip_non_ascii)
def _sanitize_messages_non_ascii(messages: list) -> bool:
"""Strip non-ASCII from a messages list in-place (ASCII-only locales); True if any stripped."""
return _sanitize_messages(messages, _strip_non_ascii, deep=False)
def _sanitize_tools_non_ascii(tools: list) -> bool:
"""Strip non-ASCII characters from tool payloads in-place."""
return _sanitize_structure_non_ascii(tools)
# In-place sanitizers; each returns True when anything changed. Surrogate repair is deep
# (tool_call ids, nested reasoning_details); the ASCII-only-locale strip is shallow.
_sanitize_structure_surrogates = partial(_sanitize_structure, fix=_sanitize_surrogates)
_sanitize_messages_surrogates = partial(_sanitize_messages, fix=_sanitize_surrogates, deep=True)
_sanitize_structure_non_ascii = partial(_sanitize_structure, fix=_strip_non_ascii)
_sanitize_messages_non_ascii = partial(_sanitize_messages, fix=_strip_non_ascii, deep=False)
_sanitize_tools_non_ascii = _sanitize_structure_non_ascii
def _escape_invalid_chars_in_json_strings(raw: str) -> str:
@@ -347,6 +325,13 @@ def _tc_field(tc: Any, key: str) -> Any:
return tc.get(key) if isinstance(tc, dict) else getattr(tc, key, None)
def _tc_set(tc: Any, key: str, value: Any) -> None:
if isinstance(tc, dict):
tc[key] = value
else:
setattr(tc, key, value)
def deterministic_call_id(fn_name: str, arguments: str, index: int = 0) -> str:
"""Deterministic call_id fallback when the API omits one (random ids would break caching)."""
seed = f"{fn_name}:{arguments}:{index}"
@@ -419,19 +404,12 @@ def uniquify_tool_call_ids(tool_calls: list) -> list:
def _renamed(value):
# Keep a composite id's response-item half so the provider's fc_/item id survives.
if isinstance(value, str) and "|" in value:
return f"{new_id}|{value.split('|', 1)[1]}"
return new_id
return f"{new_id}|{value.split('|', 1)[1]}" if isinstance(value, str) and "|" in value else new_id
try:
if isinstance(tc, dict):
tc["id"] = _renamed(tc["id"]) if tc.get("id") else new_id
if tc.get("call_id"):
tc["call_id"] = new_id
else:
tc.id = _renamed(getattr(tc, "id", None))
if getattr(tc, "call_id", None):
tc.call_id = new_id
_tc_set(tc, "id", _renamed(_tc_field(tc, "id")))
if _tc_field(tc, "call_id"):
_tc_set(tc, "call_id", new_id)
except Exception:
logger.warning("Could not uniquify duplicate tool call id %s", cid)
continue
+25 -38
View File
@@ -43,26 +43,27 @@ class MicroCompactionMixin:
"""
if head_end < self._micro_compact_cursor < tail_start:
return self._micro_compact_cursor
last_summary_idx = -1
for idx in range(head_end, tail_start):
if self._is_context_summary_message(messages[idx]):
last_summary_idx = idx
last_summary_idx = max(
(idx for idx in range(head_end, tail_start) if self._is_context_summary_message(messages[idx])),
default=-1,
)
cursor = head_end
if last_summary_idx >= head_end:
cursor = last_summary_idx + 1
# Resumed session: rehydrate the rolling summary from the surviving marker so the next
# pass merges, not replaces.
if not self._micro_compact_rolling_summary.strip():
recovered = self._rolling_summary_from_marker(messages[last_summary_idx].get("content"))
if recovered:
self._micro_compact_rolling_summary = recovered
# Rehydration proves containment: this marker (batch or micro) becomes
# supersede/defrag-eligible; unabsorbed markers never get the key.
messages[last_summary_idx][_cc().MICRO_COMPACT_MARKER_KEY] = True
logger.info(
"Micro-compaction: recovered rolling summary from "
"transcript (%d chars)", len(recovered),
)
recovered = "" if self._micro_compact_rolling_summary.strip() else (
self._rolling_summary_from_marker(messages[last_summary_idx].get("content"))
)
if recovered:
self._micro_compact_rolling_summary = recovered
# Rehydration proves containment: this marker (batch or micro) becomes
# supersede/defrag-eligible; unabsorbed markers never get the key.
messages[last_summary_idx][_cc().MICRO_COMPACT_MARKER_KEY] = True
logger.info(
"Micro-compaction: recovered rolling summary from "
"transcript (%d chars)", len(recovered),
)
self._micro_compact_cursor = cursor
return cursor
@@ -94,17 +95,11 @@ class MicroCompactionMixin:
# Boundary must close the turn: a mid-turn stop at tail_start would put the assistant marker
# beside assistant/tool rows. Any other role is a safe splice (avoids wedging the cursor).
if idx >= len(messages):
return None
boundary = messages[idx]
boundary = messages[idx] if idx < len(messages) else None
if not isinstance(boundary, dict) or boundary.get("role") in ("assistant", "tool"):
return None
return (exchange_start, idx)
def _serialize_one_exchange(self, messages: List[Dict[str, Any]], start: int, end: int) -> str:
"""Serialize a single exchange for the micro-summarizer via ``_serialize_for_summary``."""
return self._serialize_for_summary(messages[start:end])
def _build_micro_summary_prompt(self, existing_summary: str, exchange_text: str) -> List[Dict[str, str]]:
"""Build the prompt messages for a single-exchange micro-summary."""
summary_block = existing_summary if existing_summary.strip() else "(No previous summary yet.)"
@@ -167,9 +162,7 @@ class MicroCompactionMixin:
message = response.choices[0].message
content = message.get("content") if isinstance(message, dict) else getattr(message, "content", message)
if not isinstance(content, str):
content = str(content) if content else ""
content = content.strip()
content = (content if isinstance(content, str) else str(content) if content else "").strip()
if not content:
logger.info("micro-summarization returned empty content")
return None
@@ -269,7 +262,7 @@ class MicroCompactionMixin:
# Cumulative iff it subsumes an earlier marker; captured before summarizing.
_cumulative = bool(self._micro_compact_rolling_summary.strip())
exchange_text = self._serialize_one_exchange(messages, exchange_start, exchange_end)
exchange_text = self._serialize_for_summary(messages[exchange_start:exchange_end])
_exchange_tokens = estimate_tokens_rough(exchange_text)
updated_summary = self._micro_summarize_one(exchange_text)
if updated_summary is None:
@@ -402,8 +395,7 @@ class MicroCompactionMixin:
Without this the old exchange rows stay ``active=1`` and a resume double-loads
both the summary and the originals.
"""
session_db = getattr(self, "_session_db", None)
session_id = getattr(self, "_session_id", "")
session_db, session_id = getattr(self, "_session_db", None), getattr(self, "_session_id", "")
if not session_db or not session_id:
return
try:
@@ -445,8 +437,7 @@ class MicroCompactionMixin:
if supersede:
marker_idxs = [i for i, m in enumerate(result) if _is_micro_marker(m)]
if len(marker_idxs) > 1:
superseded = set(marker_idxs[:-1])
result = self._merge_adjacent_user_turns([m for i, m in enumerate(result) if i not in superseded])
result = self._merge_adjacent_user_turns([m for i, m in enumerate(result) if i not in marker_idxs[:-1]])
# Deliberately no _strip_persistence_markers: micro archives in place under the same session
# id, so stamps stay accurate and a failed archive keeps the append-only flush idempotent.
@@ -477,12 +468,8 @@ class MicroCompactionMixin:
for msg in result:
prev = merged[-1] if merged else None
if _plain_user(msg) and _plain_user(prev):
prev_content, new_content = prev["content"], msg["content"]
prev["content"] = (
(prev_content + "\n\n" + new_content) if prev_content and new_content else (prev_content or new_content)
)
# Merged content invalidates the api_content sidecar.
drop_stale_api_content(prev)
continue
merged.append(msg)
prev["content"] = "\n\n".join(c for c in (prev["content"], msg["content"]) if c)
drop_stale_api_content(prev) # merged content invalidates the api_content sidecar
else:
merged.append(msg)
return merged
+27 -42
View File
@@ -159,35 +159,30 @@ def _extract_item_text(item: Any) -> Optional[str]:
if isinstance(content, str):
return content if content.strip() else None
if isinstance(content, list):
parts = []
for part in content:
if isinstance(part, str):
candidates = (part,)
elif isinstance(part, dict):
part_meta = part.get("metadata")
candidates = (
part.get("text") or part.get("input_text") or part.get("output_text"),
part_meta.get("text") if isinstance(part_meta, dict) else None,
)
else:
continue
parts.extend(c.strip() for c in candidates if isinstance(c, str) and c.strip())
text = " ".join(parts)
return text if text.strip() else None
return None
if not isinstance(content, list):
return None
parts = []
for part in content:
if isinstance(part, str):
candidates = (part,)
elif isinstance(part, dict):
part_meta = part.get("metadata")
candidates = (
part.get("text") or part.get("input_text") or part.get("output_text"),
part_meta.get("text") if isinstance(part_meta, dict) else None,
)
else:
continue
parts.extend(c.strip() for c in candidates if isinstance(c, str) and c.strip())
text = " ".join(parts)
return text if text.strip() else None
def _has_retainable_image_content(item: Any) -> bool:
"""True for a converted Responses message with a valid ``input_image`` part (only the
adapter-owned shape counts, so empty multipart placeholders never become durable history)."""
if not isinstance(item, dict):
return False
content = item.get("content")
if not isinstance(content, list):
return False
return any(
content = item.get("content") if isinstance(item, dict) else None
return isinstance(content, list) and any(
isinstance(part, dict)
and str(part.get("type") or "").strip().lower() == "input_image"
and isinstance(part.get("image_url"), str)
@@ -234,22 +229,14 @@ def prune_pre_checkpoint_items(
"""
if not isinstance(items, list) or not items:
return items
last_cp = None
for i, item in enumerate(items):
if _is_compaction_item(item):
last_cp = i
last_cp = max((i for i, item in enumerate(items) if _is_compaction_item(item)), default=None)
if last_cp is None:
return items
first_cp = last_cp
while first_cp > 0 and _is_compaction_item(items[first_cp - 1]):
first_cp -= 1
pre = items[:first_cp]
checkpoint_run = items[first_cp : last_cp + 1]
post = items[last_cp + 1 :]
has_sources = isinstance(item_sources, list) and len(item_sources) == len(items)
pre_sources: List[Any] = item_sources[:first_cp] if has_sources else [None] * len(pre)
@@ -309,13 +296,12 @@ def prune_pre_checkpoint_items(
retained_reversed.append(item)
user_remaining -= cost
elif isinstance(item.get("content"), str):
truncated = dict(item)
truncated["content"] = item["content"][: user_remaining * 4]
truncated = {**item, "content": item["content"][: user_remaining * 4]}
if truncated["content"].strip():
retained_reversed.append(truncated)
user_remaining = 0
result = checkpoint_run + list(reversed(retained_reversed)) + post
result = items[first_cp : last_cp + 1] + list(reversed(retained_reversed)) + items[last_cp + 1 :]
logger.debug(
"Pruned pre-checkpoint items: %d input -> %d retained (user_rem=%d, summary_rem=%d)",
@@ -342,12 +328,11 @@ def is_native_compaction_rejection(error: Any, status_code: Any = None) -> bool:
text = str(error or "").lower()
if "context_management" not in text and "compact_threshold" not in text:
return False
if status_code is not None:
try:
if int(status_code) != 400:
return False
except (TypeError, ValueError):
pass
try:
if status_code is not None and int(status_code) != 400:
return False
except (TypeError, ValueError):
pass
return any(marker in text for marker in _REJECTION_MARKERS)
+13 -35
View File
@@ -47,36 +47,24 @@ def _apply_cache_marker(
role = msg.get("role", "")
content = msg.get("content")
if role == "tool" and native_anthropic:
# Top-level marker; the native adapter moves it inside tool_result.
msg["cache_control"] = cache_marker
return
if role == "tool" and not tool_part_markers:
if role == "tool" and not native_anthropic and not tool_part_markers:
# LiteLLM-style envelope: a part marker → tool_result.content[0] → non-retryable 400.
return
if content is None or content == "":
# Envelope layout: OpenRouter rejects top-level cache_control on role:tool (silent
# hang) and ignores it on empty assistant turns — no content part to carry it.
if role in ("tool", "assistant") and not native_anthropic:
return
msg["cache_control"] = cache_marker
return
if isinstance(content, str):
if (role == "tool" and native_anthropic) or content is None or content == "":
# Native role:tool: top-level marker, the adapter moves it inside tool_result. Empty
# content: no part can carry it, and OpenRouter rejects a top-level marker on role:tool
# (silent hang) and ignores it on empty assistant turns — skip those on the envelope.
if not (role in ("tool", "assistant") and not native_anthropic):
msg["cache_control"] = cache_marker
elif isinstance(content, str):
stable_prefix = find_stable_prefix(content) if role == "user" else None
if stable_prefix is not None and content[len(stable_prefix):].strip():
# Builder-declared boundary: the scaffold carries the breakpoint and the volatile
# tail rides unmarked. Request-local only — the stored message stays a string.
msg["content"] = [
_text_part(stable_prefix, cache_marker),
_text_part(content[len(stable_prefix):]),
]
msg["content"] = [_text_part(stable_prefix, cache_marker), _text_part(content[len(stable_prefix):])]
else:
msg["content"] = [_text_part(content, cache_marker)]
return
if isinstance(content, list) and content and isinstance(content[-1], dict):
elif isinstance(content, list) and content and isinstance(content[-1], dict):
content[-1]["cache_control"] = cache_marker
@@ -93,17 +81,12 @@ def _can_carry_marker(msg: dict, native_anthropic: bool, tool_part_markers: bool
if msg.get("role") == "tool" and not tool_part_markers:
return False
content = msg.get("content")
if isinstance(content, list):
return bool(content) and isinstance(content[-1], dict)
return isinstance(content, str) and content != ""
return isinstance(content[-1], dict) if isinstance(content, list) and content else isinstance(content, str) and content != ""
def _build_marker(ttl: str) -> Dict[str, str]:
"""Build a cache_control marker dict for the given TTL ('5m' or '1h')."""
marker: Dict[str, str] = {"type": "ephemeral"}
if ttl == "1h":
marker["ttl"] = "1h"
return marker
return {"type": "ephemeral", "ttl": "1h"} if ttl == "1h" else {"type": "ephemeral"}
# Alibaba-family providers (Qwen routes): five-minute context cache, 1h tier rejected. Shared
@@ -160,12 +143,7 @@ def _apply_system_cache_markers(
prompt IS the prefix the whole message is one block — never an empty text block (400).
"""
content = message.get("content")
if (
isinstance(static_system_prefix, str)
and static_system_prefix
and isinstance(content, str)
and content.startswith(static_system_prefix)
):
if isinstance(static_system_prefix, str) and static_system_prefix and isinstance(content, str) and content.startswith(static_system_prefix):
suffix = content[len(static_system_prefix):]
if suffix.strip():
message["content"] = [