Compare commits

..

No commits in common. "f5c2c91057fa82f9af0648564cd6fcacc628e7e8" and "fa1f27bad5e60eb649adc2a807bef0339144c86e" have entirely different histories.

6 changed files with 129 additions and 688 deletions

View file

@ -266,22 +266,21 @@ const MnemosynePlugin: Plugin = async (ctx: PluginInput): Promise<Hooks> => {
} }
} }
// Inject baseURL for the Mnemosyne provider only. // Inject baseURL for the Anthropic provider
// Keep provider ID `anthropic` untouched so Claude Code OAuth and // This routes all Anthropic API calls through the Mnemosyne gateway
// oh-my-opencode variants continue to work against direct Anthropic.
inputConfig.provider = inputConfig.provider ?? {}; inputConfig.provider = inputConfig.provider ?? {};
inputConfig.provider.mnemosyne = inputConfig.provider.mnemosyne ?? {}; inputConfig.provider.anthropic = inputConfig.provider.anthropic ?? {};
inputConfig.provider.mnemosyne.options = inputConfig.provider.anthropic.options =
inputConfig.provider.mnemosyne.options ?? {}; inputConfig.provider.anthropic.options ?? {};
// Only set if not already overridden by user // Only set if not already overridden by user
if (!inputConfig.provider.mnemosyne.options.baseURL) { if (!inputConfig.provider.anthropic.options.baseURL) {
// Anthropic SDK appends /messages to baseURL, so we need /v1 suffix // Anthropic SDK appends /messages to baseURL, so we need /v1 suffix
inputConfig.provider.mnemosyne.options.baseURL = `${gateway.url}/v1`; inputConfig.provider.anthropic.options.baseURL = `${gateway.url}/v1`;
log.info(`Routing Mnemosyne through gateway at ${gateway.url}/v1`); log.info(`Routing Anthropic through gateway at ${gateway.url}/v1`);
} else { } else {
log.debug( log.debug(
`baseURL already set to ${inputConfig.provider.mnemosyne.options.baseURL}, not overriding` `baseURL already set to ${inputConfig.provider.anthropic.options.baseURL}, not overriding`
); );
} }
}, },

View file

@ -287,8 +287,6 @@ class TokenMetrics:
# Payload sizes # Payload sizes
incoming_bytes_per_turn: list[int] = field(default_factory=list) incoming_bytes_per_turn: list[int] = field(default_factory=list)
outgoing_bytes_per_turn: list[int] = field(default_factory=list) outgoing_bytes_per_turn: list[int] = field(default_factory=list)
# Eviction savings (bytes removed from context per turn)
bytes_saved_per_turn: list[int] = field(default_factory=list)
def record_turn( def record_turn(
self, self,
@ -298,15 +296,14 @@ class TokenMetrics:
cache_create: int, cache_create: int,
incoming_bytes: int, incoming_bytes: int,
outgoing_bytes: int, outgoing_bytes: int,
bytes_saved: int = 0,
) -> None: ) -> None:
self.turns += 1 self.turns += 1
self.input_tokens_per_turn.append(input_tokens) self.input_tokens_per_turn.append(input_tokens)
self.effective_tokens_per_turn.append(effective_tokens) self.effective_tokens_per_turn.append(effective_tokens)
self.cache_read_per_turn.append(cache_read) self.cache_read_per_turn.append(cache_read)
self.cache_create_per_turn.append(cache_create)
self.incoming_bytes_per_turn.append(incoming_bytes) self.incoming_bytes_per_turn.append(incoming_bytes)
self.outgoing_bytes_per_turn.append(outgoing_bytes) self.outgoing_bytes_per_turn.append(outgoing_bytes)
self.bytes_saved_per_turn.append(bytes_saved)
@property @property
def total_input_tokens(self) -> int: def total_input_tokens(self) -> int:
@ -325,19 +322,14 @@ class TokenMetrics:
total_cache = self.total_cache_read + sum(self.cache_create_per_turn) total_cache = self.total_cache_read + sum(self.cache_create_per_turn)
return self.total_cache_read / total_cache if total_cache > 0 else 0.0 return self.total_cache_read / total_cache if total_cache > 0 else 0.0
@property
def total_bytes_saved(self) -> int:
return sum(self.bytes_saved_per_turn)
@property @property
def context_reduction_ratio(self) -> float: def context_reduction_ratio(self) -> float:
"""Fraction of context bytes removed by eviction. """How much smaller outgoing payloads are vs incoming.
0.0 = no reduction, 0.5 = half evicted, 1.0 = fully evicted. 1.0 = no reduction, 0.5 = halved, 0.2 = 80% reduction.
""" """
total_in = sum(self.incoming_bytes_per_turn) total_in = sum(self.incoming_bytes_per_turn)
if total_in <= 0: total_out = sum(self.outgoing_bytes_per_turn)
return 0.0 return total_out / total_in if total_in > 0 else 1.0
return min(sum(self.bytes_saved_per_turn) / total_in, 1.0)
def to_dict(self) -> dict[str, Any]: def to_dict(self) -> dict[str, Any]:
return { return {
@ -477,10 +469,7 @@ class BenchmarkCollector:
"max_ms": round(max_time, 2), "max_ms": round(max_time, 2),
} }
total_bytes_saved = sum(s.tokens.total_bytes_saved for s in sessions) context_reduction = total_outgoing / total_incoming if total_incoming > 0 else 1.0
context_reduction = (
min(total_bytes_saved / total_incoming, 1.0) if total_incoming > 0 else 0.0
)
return { return {
"sessions": n, "sessions": n,

View file

@ -20,18 +20,15 @@ from pathlib import Path
@dataclass @dataclass
class BlockEntry: class BlockEntry:
"""A tracked conversation block.""" """A tracked conversation block."""
block_id: str # Short hex ID (first 8 chars of content hash)
block_id: str # Short hex ID (first 8 chars of content hash) content_hash: str # Full SHA-256 of content
content_hash: str # Full SHA-256 of content size: int # Byte size of original content
size: int # Byte size of original content turn: int # Turn when first seen
turn: int # Turn when first seen role: str # "user" or "assistant"
role: str # "user" or "assistant" preview: str # First 80 chars for logging
preview: str # First 80 chars for logging
status: str = "resident" # resident | anchored | summarized | dropped status: str = "resident" # resident | anchored | summarized | dropped
original_content: str | None = None # Full content for fault restoration original_content: str | None = None # Full content for fault restoration
summary: str | None = None # Model-authored summary (if summarized) summary: str | None = None # Model-authored summary (if summarized)
collapse_start_turn: int | None = None
collapse_end_turn: int | None = None
class BlockStore: class BlockStore:
@ -67,22 +64,11 @@ class BlockStore:
its content hash. Labels are stable across turns as long as its content hash. Labels are stable across turns as long as
the content doesn't change. the content doesn't change.
Turn numbers are derived from message position: each user/assistant
pair is one conversation turn (user msg at index i → turn i//2 + 1).
This ensures ``collapse_range(1, 72)`` targets the right messages.
Only labels user and assistant text messages. Tool_use and Only labels user and assistant text messages. Tool_use and
tool_result blocks are managed by the PageStore, not here. tool_result blocks are managed by the PageStore, not here.
""" """
# Compute per-message turn based on position in conversation for msg in messages:
turn_counter = 0
for i, msg in enumerate(messages):
role = msg.get("role", "") role = msg.get("role", "")
# Increment turn on each user message (user+assistant = 1 turn)
if role == "user":
turn_counter += 1
msg_turn = turn_counter if turn_counter > 0 else 1
if role not in ("user", "assistant"): if role not in ("user", "assistant"):
continue continue
@ -92,14 +78,12 @@ class BlockStore:
if isinstance(content, str): if isinstance(content, str):
# Skip if already labeled by us (validated against known IDs) # Skip if already labeled by us (validated against known IDs)
if self._has_our_label(content): if self._has_our_label(content):
# Update turn on existing block if it changed
self._update_turn(content, msg_turn)
continue continue
# Skip very short messages (not worth labeling) # Skip very short messages (not worth labeling)
if len(content) < 200: if len(content) < 200:
continue continue
entry = self._get_or_create(content, role, msg_turn) entry = self._get_or_create(content, role, current_turn)
if entry and entry.status == "resident": if entry and entry.status == "resident":
size_k = entry.size / 1024 size_k = entry.size / 1024
msg["content"] = f"[tensor:{entry.block_id} ({size_k:.1f}KB)]\n{content}" msg["content"] = f"[tensor:{entry.block_id} ({size_k:.1f}KB)]\n{content}"
@ -114,28 +98,18 @@ class BlockStore:
text = block.get("text", "") text = block.get("text", "")
# Skip if already labeled by us (validated against known IDs) # Skip if already labeled by us (validated against known IDs)
if self._has_our_label(text): if self._has_our_label(text):
self._update_turn(text, msg_turn)
continue continue
# Skip short blocks # Skip short blocks
if len(text) < 200: if len(text) < 200:
continue continue
entry = self._get_or_create(text, role, msg_turn) entry = self._get_or_create(text, role, current_turn)
if entry and entry.status == "resident": if entry and entry.status == "resident":
size_k = entry.size / 1024 size_k = entry.size / 1024
block["text"] = f"[tensor:{entry.block_id} ({size_k:.1f}KB)]\n{text}" block["text"] = f"[tensor:{entry.block_id} ({size_k:.1f}KB)]\n{text}"
def _update_turn(self, labeled_content: str, turn: int) -> None: def _get_or_create(self, content: str, role: str,
"""Update the turn number on an already-labeled block.""" turn: int) -> BlockEntry | None:
import re
m = re.match(r"\[tensor:([a-f0-9]{8,12})", labeled_content)
if m:
entry = self._by_id.get(m.group(1))
if entry:
entry.turn = turn
def _get_or_create(self, content: str, role: str, turn: int) -> BlockEntry | None:
"""Get existing entry by content hash, or create a new one.""" """Get existing entry by content hash, or create a new one."""
content_hash = hashlib.sha256(content.encode()).hexdigest() content_hash = hashlib.sha256(content.encode()).hexdigest()
short_id = content_hash[:8] short_id = content_hash[:8]
@ -207,7 +181,8 @@ class BlockStore:
entry.status = "anchored" entry.status = "anchored"
return True return True
def collapse_range(self, start_turn: int, end_turn: int, summary: str) -> list[str]: def collapse_range(self, start_turn: int, end_turn: int,
summary: str) -> list[str]:
"""Replace all blocks in a turn range with a summary marker. """Replace all blocks in a turn range with a summary marker.
Marks all resident/anchored blocks in [start_turn, end_turn] as Marks all resident/anchored blocks in [start_turn, end_turn] as
@ -219,7 +194,8 @@ class BlockStore:
""" """
collapsed_ids = [] collapsed_ids = []
for entry in self._by_id.values(): for entry in self._by_id.values():
if start_turn <= entry.turn <= end_turn and entry.status in ("resident", "anchored"): if (start_turn <= entry.turn <= end_turn
and entry.status in ("resident", "anchored")):
entry.status = "dropped" entry.status = "dropped"
collapsed_ids.append(entry.block_id) collapsed_ids.append(entry.block_id)
@ -227,7 +203,9 @@ class BlockStore:
return [] return []
# Create a synthetic summary block for the range # Create a synthetic summary block for the range
synthetic_content = f"[Turns {start_turn}-{end_turn} collapsed: {summary}]" synthetic_content = (
f"[Turns {start_turn}-{end_turn} collapsed: {summary}]"
)
content_hash = hashlib.sha256(synthetic_content.encode()).hexdigest() content_hash = hashlib.sha256(synthetic_content.encode()).hexdigest()
short_id = content_hash[:8] short_id = content_hash[:8]
@ -245,8 +223,6 @@ class BlockStore:
preview=synthetic_content[:80], preview=synthetic_content[:80],
status="summarized", status="summarized",
summary=summary, summary=summary,
collapse_start_turn=start_turn,
collapse_end_turn=end_turn,
) )
self._by_id[short_id] = entry self._by_id[short_id] = entry
self._by_hash[content_hash] = short_id self._by_hash[content_hash] = short_id
@ -261,67 +237,23 @@ class BlockStore:
Modifies messages in-place. Returns stats dict. Modifies messages in-place. Returns stats dict.
""" """
stats = {"dropped": 0, "summarized": 0, "anchored": 0} stats = {"dropped": 0, "summarized": 0, "anchored": 0}
emitted_collapses: set[str] = set()
filtered_messages: list[dict] = []
for msg in messages: for msg in messages:
content = msg.get("content", "") content = msg.get("content", "")
if isinstance(content, str): if isinstance(content, str):
new_text = self._apply_to_text(content, stats, emitted_collapses) msg["content"] = self._apply_to_text(content, msg, stats)
if not new_text.strip():
continue
msg["content"] = new_text
filtered_messages.append(msg)
elif isinstance(content, list): elif isinstance(content, list):
new_blocks = []
for block in content: for block in content:
if not isinstance(block, dict): if not isinstance(block, dict) or block.get("type") != "text":
new_blocks.append(block)
continue continue
if block.get("type") != "text":
new_blocks.append(block)
continue
text = block.get("text", "") text = block.get("text", "")
new_text = self._apply_to_text(text, stats, emitted_collapses) block["text"] = self._apply_to_text(text, msg, stats)
if not new_text.strip():
continue
block["text"] = new_text
new_blocks.append(block)
if not new_blocks:
continue
msg["content"] = new_blocks
filtered_messages.append(msg)
else:
filtered_messages.append(msg)
messages[:] = filtered_messages
return stats return stats
def _find_collapse_summary(self, turn: int) -> BlockEntry | None: def _apply_to_text(self, text: str, msg: dict, stats: dict) -> str:
"""Return the newest synthetic collapse summary covering this turn."""
for entry in reversed(list(self._by_id.values())):
if (
entry.status == "summarized"
and entry.summary
and entry.collapse_start_turn is not None
and entry.collapse_end_turn is not None
and entry.collapse_start_turn <= turn <= entry.collapse_end_turn
):
return entry
return None
def _apply_to_text(
self,
text: str,
stats: dict,
emitted_collapses: set[str],
) -> str:
"""Apply block status to a single text content.""" """Apply block status to a single text content."""
m = self._BLOCK_LABEL_RE.match(text) m = self._BLOCK_LABEL_RE.match(text)
if not m: if not m:
@ -333,28 +265,18 @@ class BlockStore:
return text return text
if entry.status == "dropped": if entry.status == "dropped":
collapse_summary = self._find_collapse_summary(entry.turn)
if collapse_summary is not None:
if collapse_summary.block_id not in emitted_collapses:
emitted_collapses.add(collapse_summary.block_id)
stats["summarized"] += 1
start_turn = collapse_summary.collapse_start_turn
end_turn = collapse_summary.collapse_end_turn
return (
f"[tensor:{collapse_summary.block_id} — summarized turns "
f"{start_turn}-{end_turn}]\n"
f"{collapse_summary.summary}"
)
stats["dropped"] += 1
return ""
stats["dropped"] += 1 stats["dropped"] += 1
turn_info = f"message {entry.turn} in session log" turn_info = f"message {entry.turn} in session log"
return f"[...archived {entry.size:,} chars, {turn_info}...]" return (
f"[...archived {entry.size:,} chars, {turn_info}...]"
)
if entry.status == "summarized" and entry.summary: if entry.status == "summarized" and entry.summary:
stats["summarized"] += 1 stats["summarized"] += 1
return f"[tensor:{block_id} — summarized, was {entry.size:,} chars]\n{entry.summary}" return (
f"[tensor:{block_id} — summarized, was {entry.size:,} chars]\n"
f"{entry.summary}"
)
# resident or anchored — no change # resident or anchored — no change
if entry.status == "anchored": if entry.status == "anchored":
@ -368,12 +290,14 @@ class BlockStore:
@property @property
def total_bytes(self) -> int: def total_bytes(self) -> int:
return sum(e.size for e in self._by_id.values() if e.status == "resident") return sum(e.size for e in self._by_id.values()
if e.status == "resident")
def large_blocks(self, min_size: int = 2000) -> list[BlockEntry]: def large_blocks(self, min_size: int = 2000) -> list[BlockEntry]:
"""Return resident blocks larger than min_size, sorted by size.""" """Return resident blocks larger than min_size, sorted by size."""
return sorted( return sorted(
[e for e in self._by_id.values() if e.status == "resident" and e.size >= min_size], [e for e in self._by_id.values()
if e.status == "resident" and e.size >= min_size],
key=lambda e: e.size, key=lambda e: e.size,
reverse=True, reverse=True,
) )
@ -401,20 +325,16 @@ class BlockStore:
""" """
entries = [] entries = []
for entry in self._by_id.values(): for entry in self._by_id.values():
entries.append( entries.append({
{ "block_id": entry.block_id,
"block_id": entry.block_id, "content_hash": entry.content_hash,
"content_hash": entry.content_hash, "size": entry.size,
"size": entry.size, "turn": entry.turn,
"turn": entry.turn, "role": entry.role,
"role": entry.role, "preview": entry.preview,
"preview": entry.preview, "status": entry.status,
"status": entry.status, "summary": entry.summary,
"summary": entry.summary, })
"collapse_start_turn": entry.collapse_start_turn,
"collapse_end_turn": entry.collapse_end_turn,
}
)
tmp = path.with_suffix(".tmp") tmp = path.with_suffix(".tmp")
tmp.write_text(json.dumps(entries, indent=2)) tmp.write_text(json.dumps(entries, indent=2))
@ -449,8 +369,6 @@ class BlockStore:
preview=rec["preview"], preview=rec["preview"],
status=rec.get("status", "resident"), status=rec.get("status", "resident"),
summary=rec.get("summary"), summary=rec.get("summary"),
collapse_start_turn=rec.get("collapse_start_turn"),
collapse_end_turn=rec.get("collapse_end_turn"),
original_content=None, original_content=None,
) )
store._by_id[entry.block_id] = entry store._by_id[entry.block_id] = entry

View file

@ -57,289 +57,8 @@ from mnemosyne.object_store import ObjectStoreBackend
from mnemosyne.pager import PageStore, compact_messages from mnemosyne.pager import PageStore, compact_messages
from mnemosyne.message_store import MessageStore from mnemosyne.message_store import MessageStore
from mnemosyne.providers import adapters from mnemosyne.providers import adapters
from mnemosyne.tags import (
parse_cleanup_tags,
parse_yuyay_response,
strip_cleanup_tags,
strip_yuyay_tags,
)
from mnemosyne.telemetry import Telemetry from mnemosyne.telemetry import Telemetry
import re as _re
# ---------------------------------------------------------------------------
# SSE Cleanup Filter — strips <memory_cleanup> / <yuyay-response> tags from
# streaming responses and executes the contained ops in real-time.
# ---------------------------------------------------------------------------
_TAG_OPEN_RE = _re.compile(r"<(memory_cleanup|yuyay-response)")
_TAG_CLOSE_RE = _re.compile(r"</(memory_cleanup|yuyay-response)>")
# Matches a trailing '<' optionally followed by a prefix of a known tag name
# or '</'. Used to detect partial tag openers that arrive split across SSE
# text deltas (e.g. "<y", "<mem", "</memory_cl").
_KNOWN_TAG_PREFIXES = ("memory_cleanup", "yuyay-response", "/memory_cleanup", "/yuyay-response")
def _has_partial_tag(buf: str) -> bool:
"""Return True if *buf* ends with a partial tag opener we care about.
Only checks the last 25 characters — a complete tag name is at most
``</yuyay-response>`` (19 chars). This prevents false positives when
the model writes prose containing '<' earlier in the buffer.
"""
# Only look at the tail of the buffer for partial tags
window = buf[-25:] if len(buf) > 25 else buf
idx = window.rfind("<")
if idx == -1:
return False
tail = window[idx + 1 :] # everything after the last '<'
if not tail:
return True # bare '<' at the very end — could be start of any tag
# If the tail contains '>' then the tag is already closed — not partial
if ">" in tail:
return False
return any(p.startswith(tail) for p in _KNOWN_TAG_PREFIXES)
class SSECleanupFilter:
"""Intercept SSE text deltas and strip/execute cleanup tags.
Anthropic SSE streams emit ``content_block_delta`` events with
``{"delta": {"type": "text_delta", "text": "..."}}``. This filter
accumulates those text fragments, detects complete
``<memory_cleanup>`` / ``<yuyay-response>`` blocks, executes the
operations against the session's BlockStore / PageStore, and rewrites
the SSE ``data:`` lines with the tags removed.
Non-text events (ping, message_start, content_block_start, etc.)
pass through unchanged.
"""
# Max text deltas to buffer while waiting for a tag to complete.
# Real tags complete within 2-3 deltas. If we exceed this, it's prose.
_MAX_BUFFERED_DELTAS = 6
def __init__(self, block_store: "BlockStore | None", page_store: "PageStore | None"):
self._bs = block_store
self._ps = page_store
self._buf = "" # accumulated text not yet flushed
self._inside_tag = False # currently buffering a tag body
self._buffered_deltas = 0 # how many deltas we've buffered without flushing
self._stats: list[str] = [] # executed ops log
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
def filter_chunk(self, raw_chunk: bytes) -> bytes:
"""Process a raw SSE chunk (may contain multiple lines).
Returns the (possibly rewritten) chunk to forward downstream.
When a text delta is suppressed (buffering inside a tag), the
preceding ``event:`` header line is also removed to avoid
producing malformed SSE (event header with no data).
"""
lines = raw_chunk.split(b"\n")
out_lines: list[bytes] = []
for line in lines:
if not line.startswith(b"data: "):
# Buffer event: lines — only emit them if the next data: line
# is kept. Non-event lines (empty lines, comments) pass through.
out_lines.append(line)
continue
json_bytes = line[6:] # strip "data: " prefix
if json_bytes.strip() in (b"", b"[DONE]"):
out_lines.append(line)
continue
try:
event = json.loads(json_bytes)
except (json.JSONDecodeError, UnicodeDecodeError):
out_lines.append(line)
continue
if event.get("type") != "content_block_delta":
out_lines.append(line)
continue
delta = event.get("delta", {})
if delta.get("type") != "text_delta":
out_lines.append(line)
continue
text = delta.get("text", "")
if not text:
out_lines.append(line)
continue
# Feed text into buffer and get cleaned output
cleaned = self._feed(text)
if cleaned is None or not cleaned:
# Suppressed — also remove the preceding "event:" line
# to avoid malformed SSE (event header with no data)
if out_lines and out_lines[-1].startswith(b"event:"):
out_lines.pop()
# Also remove trailing empty line if present
if out_lines and out_lines[-1] == b"":
out_lines.pop()
continue
elif cleaned == text:
# No change — pass through original bytes
out_lines.append(line)
else:
# Rewrite the delta with cleaned text
delta["text"] = cleaned
out_lines.append(b"data: " + json.dumps(event, ensure_ascii=False).encode("utf-8"))
return b"\n".join(out_lines)
@property
def stats(self) -> str:
return "; ".join(self._stats) if self._stats else ""
def flush(self) -> str:
"""Flush any remaining buffered text (call at stream end)."""
if self._buf:
result = self._buf
self._buf = ""
self._inside_tag = False
return result
return ""
# ------------------------------------------------------------------
# Internal
# ------------------------------------------------------------------
def _feed(self, text: str) -> str | None:
"""Feed a text delta fragment. Returns cleaned text or None to suppress.
Returns:
str — cleaned text to emit (may be empty string to skip event)
None — suppress entirely (buffering inside a tag)
"""
self._buf += text
self._buffered_deltas += 1
has_open = bool(_TAG_OPEN_RE.search(self._buf))
has_close = bool(_TAG_CLOSE_RE.search(self._buf))
partial = _has_partial_tag(self._buf)
# Safety valve: if we've buffered many deltas waiting for a *partial*
# opener to resolve (e.g. "<m" that turned out to be prose "<model>"),
# flush everything. But NEVER flush when _inside_tag is True — that
# means we matched a real opening tag and are waiting for the close.
# Real tags (e.g. long collapse summaries) can span 30+ deltas.
if (
self._buffered_deltas > self._MAX_BUFFERED_DELTAS
and not self._inside_tag
and not has_close
):
result = self._buf
self._buf = ""
self._inside_tag = False
self._buffered_deltas = 0
return result
# Fast path: no tag markers and no partial tag at the end
if not self._inside_tag and not has_open and not partial:
result = self._buf
self._buf = ""
self._buffered_deltas = 0
return result
# Partial tag opener at the end (e.g. "<y", "<mem", "</memory_cl")
# but no complete open tag yet — hold buffer to accumulate more
if not self._inside_tag and not has_open and partial:
idx = self._buf.rfind("<")
if idx > 0:
emit = self._buf[:idx]
self._buf = self._buf[idx:]
return emit
return None
# We see a complete tag opening
if not self._inside_tag and has_open:
self._inside_tag = True
m = _TAG_OPEN_RE.search(self._buf)
if m and m.start() > 0:
emit = self._buf[: m.start()]
self._buf = self._buf[m.start() :]
return emit
# Check if we have a complete tag
if self._inside_tag and has_close:
self._execute_ops(self._buf)
cleaned = strip_yuyay_tags(strip_cleanup_tags(self._buf))
self._buf = ""
self._inside_tag = False
self._buffered_deltas = 0
if _TAG_OPEN_RE.search(cleaned):
self._buf = cleaned
self._inside_tag = True
return None
if _has_partial_tag(cleaned):
idx = cleaned.rfind("<")
if idx > 0:
emit = cleaned[:idx]
self._buf = cleaned[idx:]
return emit
self._buf = cleaned
return None
return cleaned
# Still inside an incomplete tag — keep buffering
if self._inside_tag:
return None
# Shouldn't reach here, but safety
result = self._buf
self._buf = ""
self._buffered_deltas = 0
return result
def _execute_ops(self, text: str) -> None:
"""Parse and execute cleanup/yuyay ops from buffered text."""
ops_list = []
cleanup_ops = parse_cleanup_tags(text)
if not cleanup_ops.empty:
ops_list.append(cleanup_ops)
yuyay_ops = parse_yuyay_response(text)
if not yuyay_ops.empty:
ops_list.append(yuyay_ops)
for ops in ops_list:
if self._bs is not None:
for block_id in ops.drops:
if self._bs.drop(block_id):
self._stats.append(f"dropped {block_id}")
for block_id, summary in ops.summaries:
if self._bs.summarize(block_id, summary):
self._stats.append(f"summarized {block_id}")
for block_id in ops.anchors:
if self._bs.anchor(block_id):
self._stats.append(f"anchored {block_id}")
for collapse in ops.collapses:
collapsed = self._bs.collapse_range(
collapse.start_turn, collapse.end_turn, collapse.summary
)
if collapsed:
self._stats.append(
f"collapsed turns {collapse.start_turn}-{collapse.end_turn} "
f"({len(collapsed)} blocks)"
)
if self._ps is not None and ops.releases:
for path in ops.releases:
self._ps.mark_released(path)
self._stats.append(f"released {len(ops.releases)} path(s)")
# ANSI for stderr status lines # ANSI for stderr status lines
_DIM = "\033[2m" _DIM = "\033[2m"
@ -1644,40 +1363,6 @@ def create_app(
"session_tokens": total_tokens, "session_tokens": total_tokens,
} }
@app.get("/api/blocks")
async def api_blocks(session_id: str | None = None) -> dict[str, Any]:
"""Debug endpoint: expose BlockStore state per session."""
all_sessions = sessions.all()
out: dict[str, Any] = {}
for sid, sess in all_sessions.items():
if session_id and sid != session_id:
continue
bs = sess.block_store
blocks = []
for bid, entry in bs._by_id.items():
blocks.append(
{
"id": bid,
"status": entry.status,
"turn": entry.turn,
"role": entry.role,
"size": entry.size,
"preview": entry.preview[:80] if entry.preview else "",
"summary": entry.summary,
"collapse_start": getattr(entry, "collapse_start_turn", None),
"collapse_end": getattr(entry, "collapse_end_turn", None),
}
)
out[sid] = {
"total_blocks": len(blocks),
"by_status": {},
"blocks": blocks,
}
for b in blocks:
s = b["status"]
out[sid]["by_status"][s] = out[sid]["by_status"].get(s, 0) + 1
return out
@app.get("/api/compaction-context") @app.get("/api/compaction-context")
async def api_compaction_context( async def api_compaction_context(
session_id: str, session_id: str,
@ -1743,10 +1428,9 @@ def create_app(
# ── Pre/post processing ────────────────────────────────────────── # ── Pre/post processing ──────────────────────────────────────────
def _preprocess(payload: dict, session: Session) -> tuple[dict, int]: def _preprocess(payload: dict, session: Session) -> dict:
"""Apply gateway transformations before pipeline and forwarding. """Apply gateway transformations before pipeline and forwarding.
Returns (modified_payload, bytes_saved_by_eviction).
Operates on the raw Anthropic-format payload (before normalization) Operates on the raw Anthropic-format payload (before normalization)
because system prompt injection needs access to the system field because system prompt injection needs access to the system field
directly. directly.
@ -1775,28 +1459,16 @@ def create_app(
ps = session.page_store ps = session.page_store
ms = session.message_store ms = session.message_store
_bytes_saved = 0
# 1. Ingest into MessageStore (asserts append-only, compacts) # 1. Ingest into MessageStore (asserts append-only, compacts)
ingest = ms.ingest( ingest = ms.ingest(
incoming_messages, incoming_messages,
age_threshold=4, age_threshold=4,
min_evict_size=min_evict_size, min_evict_size=min_evict_size,
) )
_bytes_saved += ingest.bytes_saved if ingest.new_count > 0 or ingest.compacted_count > 0:
if ingest.physical_tail_deleted > 0:
session._segmented_objects = [
obj
for obj in session._segmented_objects
if obj.turn_end < ingest.deleted_physical_start
]
if ingest.new_count > 0 or ingest.compacted_count > 0 or ingest.physical_tail_deleted > 0:
parts = [] parts = []
if ingest.new_count: if ingest.new_count:
parts.append(f"+{ingest.new_count} msgs") parts.append(f"+{ingest.new_count} msgs")
if ingest.physical_tail_deleted:
parts.append(f"undo removed {ingest.physical_tail_deleted} msgs")
if ingest.compacted_count: if ingest.compacted_count:
parts.append(f"{ingest.compacted_count} evicted") parts.append(f"{ingest.compacted_count} evicted")
if ingest.bytes_saved: if ingest.bytes_saved:
@ -1810,21 +1482,17 @@ def create_app(
file=sys.stderr, file=sys.stderr,
) )
# 1b. Segment only newly ingested physical messages into semantic objects # 1b. Segment new messages into semantic objects and store in ObjectStore
# and store them in ObjectStore. Re-segmenting the full history here # Phase 4d: admission control gates each object before storage
# would balloon object counts across turns.
try: try:
if ingest.new_count > 0: with Timer(session.benchmark.latency["segmentation"]):
new_physical_messages = ms.messages[ segmented = session.segmenter.segment_incremental(
ingest.new_physical_start : ingest.new_physical_start + ingest.new_count ms.messages,
] session._segmented_objects,
with Timer(session.benchmark.latency["segmentation"]): start_turn=0,
segmented = session.segmenter.segment_incremental( )
new_physical_messages, new_count = len(segmented) - len(session._segmented_objects)
session._segmented_objects, if new_count > 0:
start_turn=ingest.new_physical_start,
)
new_count = len(segmented) - len(session._segmented_objects)
new_objects = segmented[len(session._segmented_objects) :] new_objects = segmented[len(session._segmented_objects) :]
admitted_count = 0 admitted_count = 0
rejected_count = 0 rejected_count = 0
@ -2095,18 +1763,6 @@ def create_app(
# 4. Build ephemeral outbound view — never mutate the physical store # 4. Build ephemeral outbound view — never mutate the physical store
payload["messages"] = copy.deepcopy(ms.messages) payload["messages"] = copy.deepcopy(ms.messages)
# Apply cleanup/block-state rewrites to the outbound view so model-authored
# drop/summarize/collapse operations actually affect future forwarded context.
cleanup_apply = session.block_store.apply_to_messages(payload["messages"])
if any(cleanup_apply.values()):
print(
f" {_DIM}[{session.id}] outbound cleanup apply: "
f"drop={cleanup_apply['dropped']} "
f"sum={cleanup_apply['summarized']} "
f"anchor={cleanup_apply['anchored']}{_RESET}",
file=sys.stderr,
)
# 4a. Phantom tool injection — DISABLED when proxying for opencode. # 4a. Phantom tool injection — DISABLED when proxying for opencode.
# opencode validates tool calls against its own registry and rejects # opencode validates tool calls against its own registry and rejects
# unknown tools like memory_query. Phantom tools require SSE stream # unknown tools like memory_query. Phantom tools require SSE stream
@ -2153,7 +1809,13 @@ def create_app(
with open(session._page_checkpoint, "w") as f: with open(session._page_checkpoint, "w") as f:
_json.dump(session.page_store.checkpoint(), f) _json.dump(session.page_store.checkpoint(), f)
return payload, _bytes_saved return payload
def _check_token_cap(usage: dict, session: Session) -> None:
"""Track usage and enforce token cap."""
session.track_usage(usage)
if token_cap <= 0:
return
effective = session.token_state["last_effective"] effective = session.token_state["last_effective"]
pct = effective / token_cap * 100 pct = effective / token_cap * 100
@ -2171,8 +1833,8 @@ def create_app(
"""Update FidelityManager window understanding from actual API usage. """Update FidelityManager window understanding from actual API usage.
After receiving the response, we know the real input token count. After receiving the response, we know the real input token count.
Scale the FM's window_size so its internal pressure calculation Update the fidelity manager's window_size understanding and schedule
reflects the real API token usage, then trigger degradation. degradation if pressure is above NORMAL for the next turn.
""" """
input_tokens = usage.get("input_tokens", 0) input_tokens = usage.get("input_tokens", 0)
if input_tokens <= 0: if input_tokens <= 0:
@ -2181,22 +1843,14 @@ def create_app(
fm = session.fidelity_manager fm = session.fidelity_manager
turn = session.token_state.get("turn", 0) turn = session.token_state.get("turn", 0)
# The FidelityManager tracks its own token budget via registered objects.
# Here we use the real API token count to check if we need proactive degradation.
# If real usage exceeds the fidelity window threshold, trigger degradation now
# so the NEXT turn benefits from reduced content.
pressure_ratio = input_tokens / fm.window_size if fm.window_size > 0 else 1.0 pressure_ratio = input_tokens / fm.window_size if fm.window_size > 0 else 1.0
if pressure_ratio >= fm.threshold_caution: if pressure_ratio >= fm.threshold_caution:
# The FM's internal pressure uses total_tokens()/window_size, but transitions = fm.degrade(turn)
# total_tokens() only counts registered objects (a fraction of the
# real context). Temporarily scale window_size down so the FM's
# pressure matches the real API pressure, then restore it.
obj_tokens = fm.total_tokens()
if obj_tokens > 0:
# Set window_size so obj_tokens/window_size == pressure_ratio
saved_ws = fm.window_size
fm.window_size = max(1, int(obj_tokens / pressure_ratio))
transitions = fm.degrade(turn)
fm.window_size = saved_ws
else:
transitions = fm.degrade(turn)
if transitions: if transitions:
zone = fm.current_pressure() zone = fm.current_pressure()
for _obj_id, old_level, new_level in transitions: for _obj_id, old_level, new_level in transitions:
@ -2227,28 +1881,11 @@ def create_app(
exc_info=True, exc_info=True,
) )
def _check_token_cap(usage: dict, session: Session) -> None:
"""Track usage and enforce token cap."""
session.track_usage(usage)
if token_cap <= 0:
return
effective = session.token_state["last_effective"]
pct = effective / token_cap * 100
sid = session.id
if effective > token_cap:
session.token_state["blocked"] = True
print(
f"{_RED} [{sid}] TOKEN CAP EXCEEDED: {effective:,} / {token_cap:,} "
f"({pct:.0f}%) — next request will be blocked{_RESET}",
file=sys.stderr,
)
def _display_turn_status( def _display_turn_status(
usage: dict, usage: dict,
session: Session, session: Session,
incoming_bytes: int = 0, incoming_bytes: int = 0,
outgoing_bytes: int = 0, outgoing_bytes: int = 0,
bytes_saved: int = 0,
) -> None: ) -> None:
"""Post-response status line with cache hit rate.""" """Post-response status line with cache hit rate."""
sid = session.id sid = session.id
@ -2268,7 +1905,6 @@ def create_app(
cache_create=cache_create, cache_create=cache_create,
incoming_bytes=incoming_bytes, incoming_bytes=incoming_bytes,
outgoing_bytes=outgoing_bytes, outgoing_bytes=outgoing_bytes,
bytes_saved=bytes_saved,
) )
cap_str = "" cap_str = ""
@ -2334,13 +1970,9 @@ def create_app(
# Cleanup now runs inside _preprocess (before manifest injection) # Cleanup now runs inside _preprocess (before manifest injection)
# Measure raw incoming size BEFORE any preprocessing
incoming_bytes = len(json.dumps(payload, default=str).encode("utf-8"))
# Pre-process: system status, block labeling # Pre-process: system status, block labeling
preprocess_bytes_saved = 0
if endpoint == "messages": if endpoint == "messages":
payload, preprocess_bytes_saved = _preprocess(payload, session) payload = _preprocess(payload, session)
# Optional provider-level model override for cost-controlled runs. # Optional provider-level model override for cost-controlled runs.
if provider == "anthropic" and anthropic_model_override: if provider == "anthropic" and anthropic_model_override:
@ -2353,6 +1985,7 @@ def create_app(
request_id = str(uuid.uuid4()) request_id = str(uuid.uuid4())
started = time.perf_counter() started = time.perf_counter()
incoming_bytes = len(json.dumps(payload, default=str).encode("utf-8"))
session_id = session.id session_id = session.id
req = adapter.normalize_request(payload) req = adapter.normalize_request(payload)
@ -2412,10 +2045,6 @@ def create_app(
bytes_out = len(body) bytes_out = len(body)
yield body yield body
else: else:
cleanup_filter = SSECleanupFilter(
block_store=session.block_store,
page_store=session.page_store,
)
for chunk in resp.iter_bytes(): for chunk in resp.iter_bytes():
bytes_out += len(chunk) bytes_out += len(chunk)
chunk_count += 1 chunk_count += 1
@ -2428,9 +2057,7 @@ def create_app(
provider=provider, provider=provider,
usage_accumulator=usage, usage_accumulator=usage,
) )
filtered = cleanup_filter.filter_chunk(chunk) yield chunk
if filtered:
yield filtered
except Exception as e: except Exception as e:
stream_error = True stream_error = True
emit_event( emit_event(
@ -2452,7 +2079,6 @@ def create_app(
session, session,
incoming_bytes=incoming_bytes, incoming_bytes=incoming_bytes,
outgoing_bytes=outgoing_bytes, outgoing_bytes=outgoing_bytes,
bytes_saved=max(0, incoming_bytes - outgoing_bytes),
) )
_update_fidelity_pressure(usage, session) _update_fidelity_pressure(usage, session)
emit_event( emit_event(
@ -2514,7 +2140,6 @@ def create_app(
session, session,
incoming_bytes=incoming_bytes, incoming_bytes=incoming_bytes,
outgoing_bytes=outgoing_bytes, outgoing_bytes=outgoing_bytes,
bytes_saved=max(0, incoming_bytes - outgoing_bytes),
) )
_update_fidelity_pressure(usage, session) _update_fidelity_pressure(usage, session)

View file

@ -79,13 +79,9 @@ def _fingerprint(msg: dict) -> str:
@dataclass @dataclass
class IngestResult: class IngestResult:
"""Result of ingesting a new turn's messages.""" """Result of ingesting a new turn's messages."""
new_count: int = 0 new_count: int = 0
new_physical_start: int = 0
mutations_detected: int = 0 mutations_detected: int = 0
deletions_detected: int = 0 deletions_detected: int = 0
physical_tail_deleted: int = 0
deleted_physical_start: int = 0
compacted_count: int = 0 compacted_count: int = 0
bytes_saved: int = 0 bytes_saved: int = 0
@ -93,7 +89,8 @@ class IngestResult:
class MessageStore: class MessageStore:
"""Pichay's compacted conversation history for a session.""" """Pichay's compacted conversation history for a session."""
def __init__(self, session_id: str, page_store: PageStore, log_path: Path | None = None): def __init__(self, session_id: str, page_store: PageStore,
log_path: Path | None = None):
self.session_id = session_id self.session_id = session_id
self.page_store = page_store self.page_store = page_store
self.log_path = log_path self.log_path = log_path
@ -127,16 +124,10 @@ class MessageStore:
return json.dumps(content, default=str)[:limit] return json.dumps(content, default=str)[:limit]
return str(content)[:limit] return str(content)[:limit]
def _log_violation( def _log_violation(self, kind: str, index: int, msg: dict | None,
self, expected_fp: str, actual_fp: str,
kind: str, old_msg: dict | None = None,
index: int, deleted_msgs: list[dict] | None = None) -> None:
msg: dict | None,
expected_fp: str,
actual_fp: str,
old_msg: dict | None = None,
deleted_msgs: list[dict] | None = None,
) -> None:
"""Log append-only violations to file for later analysis.""" """Log append-only violations to file for later analysis."""
if self.log_path is None: if self.log_path is None:
return return
@ -197,26 +188,20 @@ class MessageStore:
self._turn += 1 self._turn += 1
result = IngestResult() result = IngestResult()
client_known = len(self._client_fps) client_known = len(self._client_fps)
first_mutation_index: int | None = None
# ── Detect mutations in known client messages ──────────── # ── Detect mutations in known client messages ────────────
check_limit = min(client_known, len(incoming)) check_limit = min(client_known, len(incoming))
for i in range(check_limit): for i in range(check_limit):
fp = _fingerprint(incoming[i]) fp = _fingerprint(incoming[i])
if fp != self._client_fps[i]: if fp != self._client_fps[i]:
if first_mutation_index is None:
first_mutation_index = i
result.mutations_detected += 1 result.mutations_detected += 1
self.total_mutations += 1 self.total_mutations += 1
# Look up physical message via mapping for comparison # Look up physical message via mapping for comparison
phys_idx = self._client_to_physical[i] phys_idx = self._client_to_physical[i]
old_msg = self._messages[phys_idx] if phys_idx < len(self._messages) else None old_msg = self._messages[phys_idx] if phys_idx < len(self._messages) else None
self._log_violation( self._log_violation(
"mutation", "mutation", i, incoming[i],
i, self._client_fps[i], fp,
incoming[i],
self._client_fps[i],
fp,
old_msg=old_msg, old_msg=old_msg,
) )
print( print(
@ -225,28 +210,8 @@ class MessageStore:
f"got {fp[:32]}{_RESET}", f"got {fp[:32]}{_RESET}",
file=sys.stderr, file=sys.stderr,
) )
# Update CLIENT fingerprint only — physical store unchanged
# Tail mutation: rebuild physical/client history from first changed index onward. self._client_fps[i] = fp
if first_mutation_index is not None:
if first_mutation_index < len(self._client_to_physical):
physical_start = self._client_to_physical[first_mutation_index]
else:
physical_start = len(self._messages)
physical_start = max(0, min(physical_start, len(self._messages)))
removed = len(self._messages) - physical_start
self._messages = self._messages[:physical_start]
self._fingerprints = self._fingerprints[:physical_start]
self._client_fps = self._client_fps[:first_mutation_index]
self._client_to_physical = self._client_to_physical[:first_mutation_index]
client_known = len(self._client_fps)
result.physical_tail_deleted = max(result.physical_tail_deleted, removed)
result.deleted_physical_start = physical_start
print(
f" {_DIM}[{self.session_id}] CLIENT TAIL MUTATION APPLIED: "
f"rebuilt history from client index {first_mutation_index} "
f"(removed {removed} physical msgs){_RESET}",
file=sys.stderr,
)
# ── Detect client deletions (compaction) ───────────────── # ── Detect client deletions (compaction) ─────────────────
if len(incoming) < client_known: if len(incoming) < client_known:
@ -254,7 +219,6 @@ class MessageStore:
result.deletions_detected = deleted result.deletions_detected = deleted
self.total_deletions += deleted self.total_deletions += deleted
self.total_client_deletions_absorbed += deleted self.total_client_deletions_absorbed += deleted
undo_applied = False
# Log what the client is deleting (from physical store via mapping) # Log what the client is deleting (from physical store via mapping)
deleted_physical = [] deleted_physical = []
for ci in range(len(incoming), client_known): for ci in range(len(incoming), client_known):
@ -262,49 +226,24 @@ class MessageStore:
if pi < len(self._messages): if pi < len(self._messages):
deleted_physical.append(self._messages[pi]) deleted_physical.append(self._messages[pi])
self._log_violation( self._log_violation(
"deletion", "deletion", client_known, None,
client_known, f"expected_{client_known}", f"got_{len(incoming)}",
None,
f"expected_{client_known}",
f"got_{len(incoming)}",
deleted_msgs=deleted_physical, deleted_msgs=deleted_physical,
) )
deleted_indices = self._client_to_physical[len(incoming) : client_known] print(
if deleted_indices: f" {_DIM}[{self.session_id}] CLIENT DELETION ABSORBED: "
valid_deleted = sorted( f"{deleted} messages dropped by client, "
{pi for pi in deleted_indices if 0 <= pi < len(self._messages)} f"physical store unchanged ({len(self._messages)} msgs){_RESET}",
) file=sys.stderr,
if valid_deleted: )
tail_start = valid_deleted[0]
removed = len(self._messages) - tail_start
self._messages = self._messages[:tail_start]
self._fingerprints = self._fingerprints[:tail_start]
result.physical_tail_deleted = removed
result.deleted_physical_start = tail_start
undo_applied = True
print(
f" {_DIM}[{self.session_id}] CLIENT UNDO APPLIED: "
f"removed {removed} tail messages from physical store{_RESET}",
file=sys.stderr,
)
if not undo_applied:
print(
f" {_DIM}[{self.session_id}] CLIENT DELETION ABSORBED: "
f"{deleted} messages dropped by client, "
f"physical store unchanged ({len(self._messages)} msgs){_RESET}",
file=sys.stderr,
)
# Truncate CLIENT tracking only — physical store stays intact # Truncate CLIENT tracking only — physical store stays intact
self._client_fps = self._client_fps[: len(incoming)] self._client_fps = self._client_fps[:len(incoming)]
self._client_to_physical = self._client_to_physical[: len(incoming)] self._client_to_physical = self._client_to_physical[:len(incoming)]
# ── Extract and append new messages ────────────────────── # ── Extract and append new messages ──────────────────────
new_start = min(client_known, len(incoming)) new_start = min(client_known, len(incoming))
new_messages = incoming[new_start:] new_messages = incoming[new_start:]
result.new_count = len(new_messages) result.new_count = len(new_messages)
result.new_physical_start = len(self._messages)
if new_messages: if new_messages:
# Deep copy new messages so we own them # Deep copy new messages so we own them
@ -317,7 +256,7 @@ class MessageStore:
_strip_cache_control(msg) _strip_cache_control(msg)
# Track in physical store and client mapping # Track in physical store and client mapping
phys_start = result.new_physical_start phys_start = len(self._messages)
for j, msg in enumerate(new_messages): for j, msg in enumerate(new_messages):
fp = _fingerprint(msg) fp = _fingerprint(msg)
self._fingerprints.append(fp) self._fingerprints.append(fp)

View file

@ -22,7 +22,6 @@ from dataclasses import dataclass, field
@dataclass @dataclass
class CollapseOp: class CollapseOp:
"""A turn-range collapse: replace multiple turns with a summary.""" """A turn-range collapse: replace multiple turns with a summary."""
start_turn: int start_turn: int
end_turn: int end_turn: int
summary: str summary: str
@ -31,7 +30,6 @@ class CollapseOp:
@dataclass @dataclass
class CleanupOps: class CleanupOps:
"""Parsed cleanup operations from a <memory_cleanup> tag.""" """Parsed cleanup operations from a <memory_cleanup> tag."""
drops: list[str] = field(default_factory=list) drops: list[str] = field(default_factory=list)
summaries: list[tuple[str, str]] = field(default_factory=list) summaries: list[tuple[str, str]] = field(default_factory=list)
anchors: list[str] = field(default_factory=list) anchors: list[str] = field(default_factory=list)
@ -40,7 +38,8 @@ class CleanupOps:
@property @property
def empty(self) -> bool: def empty(self) -> bool:
return not (self.drops or self.summaries or self.anchors or self.releases or self.collapses) return not (self.drops or self.summaries or self.anchors
or self.releases or self.collapses)
def __str__(self) -> str: def __str__(self) -> str:
parts = [] parts = []
@ -66,95 +65,67 @@ _TAG_PATTERN = re.compile(
# Match tensor/block ID references: tensor:xxxxxxxx or block:xxxxxxxx (8-12 hex chars) # Match tensor/block ID references: tensor:xxxxxxxx or block:xxxxxxxx (8-12 hex chars)
_BLOCK_ID = re.compile(r"(?:tensor|block):([a-f0-9]{8,12})(?![a-f0-9])") _BLOCK_ID = re.compile(r"(?:tensor|block):([a-f0-9]{8,12})(?![a-f0-9])")
# Match summarize with quoted summary text (prose format) # Match summarize with quoted summary text
_SUMMARIZE_PATTERN = re.compile( _SUMMARIZE_PATTERN = re.compile(
r'summarize:\s*(?:tensor|block):([a-f0-9]{8,12})(?![a-f0-9])\s+"([^"]*)"' r'summarize:\s*(?:tensor|block):([a-f0-9]{8,12})(?![a-f0-9])\s+"([^"]*)"'
) )
# Match release with comma-separated paths (prose format) # Match release with comma-separated paths
_RELEASE_PATTERN = re.compile(r"release:\s*(.+)") _RELEASE_PATTERN = re.compile(r"release:\s*(.+)")
# Match collapse with turn range and quoted summary # Match collapse with turn range and quoted summary
# Format: collapse: turns 3-8 "Summary of what happened in those turns" # Format: collapse: turns 3-8 "Summary of what happened in those turns"
# Also: <collapse>turns 3-8 "Summary"</collapse> _COLLAPSE_PATTERN = re.compile(
_COLLAPSE_PATTERN = re.compile(r'collapse[\s:>]+turns\s+(\d+)\s*-\s*(\d+)\s+"([^"]*)"') r'collapse:\s*turns\s+(\d+)\s*-\s*(\d+)\s+"([^"]*)"'
)
# --- XML element variants (model sometimes uses XML instead of prose) ---
# <drop>block:4ab5ee7a</drop> or <drop>tensor:4ab5ee7a</drop>
_XML_DROP = re.compile(r"<drop>(?:tensor|block):([a-f0-9]{8,12})</drop>")
# <release handle="abc123"/> or <release handle="abc123" reason="..."/>
_XML_RELEASE = re.compile(r'<release\s+handle="([a-f0-9]{8,12})"')
# <anchor>block:abc123</anchor>
_XML_ANCHOR = re.compile(r"<anchor>(?:tensor|block):([a-f0-9]{8,12})</anchor>")
# <summarize id="abc123">summary text</summarize>
_XML_SUMMARIZE = re.compile(r'<summarize\s+id="([a-f0-9]{8,12})"[^>]*>([^<]*)</summarize>')
# <collapse>turns 3-8 "Summary"</collapse>
_XML_COLLAPSE = re.compile(r'<collapse>turns\s+(\d+)\s*-\s*(\d+)\s+"([^"]*)"</collapse>')
def parse_cleanup_tags(text: str) -> CleanupOps: def parse_cleanup_tags(text: str) -> CleanupOps:
"""Extract memory operations from ``<memory_cleanup>`` blocks. """Extract cleanup operations from text containing <memory_cleanup> tags.
Supports both prose format (``drop: block:abc``) and XML element Returns a CleanupOps with all parsed operations. Multiple tags in
format (``<drop>block:abc</drop>``). the same text are merged into a single CleanupOps.
""" """
ops = CleanupOps() ops = CleanupOps()
for match in _TAG_PATTERN.finditer(text): for match in _TAG_PATTERN.finditer(text):
body = match.group(1) body = match.group(1)
# --- XML element variants (scan full body first) ---
for m in _XML_DROP.finditer(body):
ops.drops.append(m.group(1))
for m in _XML_RELEASE.finditer(body):
ops.releases.append(m.group(1))
for m in _XML_ANCHOR.finditer(body):
ops.anchors.append(m.group(1))
for m in _XML_SUMMARIZE.finditer(body):
ops.summaries.append((m.group(1), m.group(2)))
for m in _XML_COLLAPSE.finditer(body):
ops.collapses.append(
CollapseOp(
start_turn=int(m.group(1)),
end_turn=int(m.group(2)),
summary=m.group(3),
)
)
# --- Prose line-based format ---
for line in body.splitlines(): for line in body.splitlines():
line = line.strip() line = line.strip()
if not line or line.startswith("<"): if not line:
continue continue
# Summarize (must check before drop — both start with block ID)
m = _SUMMARIZE_PATTERN.match(line) m = _SUMMARIZE_PATTERN.match(line)
if m: if m:
ops.summaries.append((m.group(1), m.group(2))) ops.summaries.append((m.group(1), m.group(2)))
continue continue
# Drop
if line.startswith("drop:"): if line.startswith("drop:"):
m = _BLOCK_ID.search(line) m = _BLOCK_ID.search(line)
if m: if m:
ops.drops.append(m.group(1)) ops.drops.append(m.group(1))
continue continue
# Anchor
if line.startswith("anchor:"): if line.startswith("anchor:"):
m = _BLOCK_ID.search(line) m = _BLOCK_ID.search(line)
if m: if m:
ops.anchors.append(m.group(1)) ops.anchors.append(m.group(1))
continue continue
# Collapse (turn range)
m = _COLLAPSE_PATTERN.match(line) m = _COLLAPSE_PATTERN.match(line)
if m: if m:
ops.collapses.append( ops.collapses.append(CollapseOp(
CollapseOp( start_turn=int(m.group(1)),
start_turn=int(m.group(1)), end_turn=int(m.group(2)),
end_turn=int(m.group(2)), summary=m.group(3),
summary=m.group(3), ))
)
)
continue continue
# Release
m = _RELEASE_PATTERN.match(line) m = _RELEASE_PATTERN.match(line)
if m: if m:
paths = [p.strip() for p in m.group(1).split(",") if p.strip()] paths = [p.strip() for p in m.group(1).split(",") if p.strip()]