fix: parse XML-format cleanup tags and strip from SSE stream

The model emits cleanup ops as XML elements (<drop>block:x</drop>,
<release handle="x"/>, <collapse>turns N-M "summary"</collapse>)
but the parser only handled prose format (drop: block:x). Add XML
regex matchers alongside the existing prose parser so both formats
are recognized, executed, and stripped from the streaming output.
This commit is contained in:
Joey Yakimowich-Payne 2026-03-13 21:23:26 -06:00
commit ad2c296ba3

View file

@ -22,6 +22,7 @@ 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
@ -30,6 +31,7 @@ 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)
@ -38,8 +40,7 @@ class CleanupOps:
@property @property
def empty(self) -> bool: def empty(self) -> bool:
return not (self.drops or self.summaries or self.anchors return not (self.drops or self.summaries or self.anchors or self.releases or self.collapses)
or self.releases or self.collapses)
def __str__(self) -> str: def __str__(self) -> str:
parts = [] parts = []
@ -65,67 +66,95 @@ _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 # Match summarize with quoted summary text (prose format)
_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 # Match release with comma-separated paths (prose format)
_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"
_COLLAPSE_PATTERN = re.compile( # Also: <collapse>turns 3-8 "Summary"</collapse>
r'collapse:\s*turns\s+(\d+)\s*-\s*(\d+)\s+"([^"]*)"' _COLLAPSE_PATTERN = re.compile(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 cleanup operations from text containing <memory_cleanup> tags. """Extract memory operations from ``<memory_cleanup>`` blocks.
Returns a CleanupOps with all parsed operations. Multiple tags in Supports both prose format (``drop: block:abc``) and XML element
the same text are merged into a single CleanupOps. format (``<drop>block:abc</drop>``).
""" """
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: if not line or line.startswith("<"):
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(CollapseOp( ops.collapses.append(
start_turn=int(m.group(1)), CollapseOp(
end_turn=int(m.group(2)), start_turn=int(m.group(1)),
summary=m.group(3), end_turn=int(m.group(2)),
)) 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()]