fix: accurate context reduction stats and SSE cleanup tag filter
Measure incoming_bytes before _preprocess() so bytes_saved reflects true reduction. Add SSECleanupFilter that intercepts memory_cleanup/yuyay-response tags in streaming responses, strips them from output, and executes ops (drops, collapses, releases) in real-time. Handles partial tags split across SSE chunks with a safety valve to flush stale buffers for prose.
This commit is contained in:
parent
2bf6baaa33
commit
92fba55f70
2 changed files with 282 additions and 5 deletions
|
|
@ -477,7 +477,10 @@ class BenchmarkCollector:
|
||||||
"max_ms": round(max_time, 2),
|
"max_ms": round(max_time, 2),
|
||||||
}
|
}
|
||||||
|
|
||||||
context_reduction = total_outgoing / total_incoming if total_incoming > 0 else 1.0
|
total_bytes_saved = sum(s.tokens.total_bytes_saved for s in sessions)
|
||||||
|
context_reduction = (
|
||||||
|
min(total_bytes_saved / total_incoming, 1.0) if total_incoming > 0 else 0.0
|
||||||
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"sessions": n,
|
"sessions": n,
|
||||||
|
|
|
||||||
|
|
@ -57,8 +57,274 @@ 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.
|
||||||
|
"""
|
||||||
|
lines = raw_chunk.split(b"\n")
|
||||||
|
out_lines: list[bytes] = []
|
||||||
|
|
||||||
|
for line in lines:
|
||||||
|
if not line.startswith(b"data: "):
|
||||||
|
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:
|
||||||
|
# Entire chunk is being buffered (inside a tag) — suppress
|
||||||
|
continue
|
||||||
|
elif cleaned == text:
|
||||||
|
# No change — pass through original bytes
|
||||||
|
out_lines.append(line)
|
||||||
|
elif not cleaned:
|
||||||
|
# Text was entirely a cleanup tag — suppress this event
|
||||||
|
continue
|
||||||
|
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 too many deltas without resolving
|
||||||
|
# a tag, this is just prose containing '<' — flush everything.
|
||||||
|
if self._buffered_deltas > self._MAX_BUFFERED_DELTAS 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"
|
||||||
|
|
@ -2011,6 +2277,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
|
preprocess_bytes_saved = 0
|
||||||
if endpoint == "messages":
|
if endpoint == "messages":
|
||||||
|
|
@ -2027,7 +2296,6 @@ 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)
|
||||||
|
|
@ -2087,6 +2355,10 @@ 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
|
||||||
|
|
@ -2099,7 +2371,9 @@ def create_app(
|
||||||
provider=provider,
|
provider=provider,
|
||||||
usage_accumulator=usage,
|
usage_accumulator=usage,
|
||||||
)
|
)
|
||||||
yield chunk
|
filtered = cleanup_filter.filter_chunk(chunk)
|
||||||
|
if filtered:
|
||||||
|
yield filtered
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
stream_error = True
|
stream_error = True
|
||||||
emit_event(
|
emit_event(
|
||||||
|
|
@ -2121,7 +2395,7 @@ def create_app(
|
||||||
session,
|
session,
|
||||||
incoming_bytes=incoming_bytes,
|
incoming_bytes=incoming_bytes,
|
||||||
outgoing_bytes=outgoing_bytes,
|
outgoing_bytes=outgoing_bytes,
|
||||||
bytes_saved=preprocess_bytes_saved if endpoint == "messages" else 0,
|
bytes_saved=max(0, incoming_bytes - outgoing_bytes),
|
||||||
)
|
)
|
||||||
_update_fidelity_pressure(usage, session)
|
_update_fidelity_pressure(usage, session)
|
||||||
emit_event(
|
emit_event(
|
||||||
|
|
@ -2183,7 +2457,7 @@ def create_app(
|
||||||
session,
|
session,
|
||||||
incoming_bytes=incoming_bytes,
|
incoming_bytes=incoming_bytes,
|
||||||
outgoing_bytes=outgoing_bytes,
|
outgoing_bytes=outgoing_bytes,
|
||||||
bytes_saved=preprocess_bytes_saved,
|
bytes_saved=max(0, incoming_bytes - outgoing_bytes),
|
||||||
)
|
)
|
||||||
_update_fidelity_pressure(usage, session)
|
_update_fidelity_pressure(usage, session)
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue