refactor(agent): finish memory/compaction/prompt-cache compaction pass (>=25% LOC)
This commit is contained in:
+38
-83
@@ -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
@@ -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 []
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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"] = [
|
||||
|
||||
Reference in New Issue
Block a user