From bee90915dbba3181a7c4942c65ee771fa46f7988 Mon Sep 17 00:00:00 2001 From: Joey Yakimowich-Payne Date: Fri, 13 Mar 2026 11:49:12 -0600 Subject: [PATCH] feat: add SSE stream filter for yuyay protocol tags MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Strip , , and tags from the SSE stream before forwarding to the client. The cooperative memory protocol tags are still processed by the gateway on the next inbound request — they just no longer leak into the user's visible output. Handles tags spanning across multiple text_delta SSE events. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- src/mnemosyne/gateway.py | 160 ++++++++++++++++++++++++++++++++++++++- 1 file changed, 159 insertions(+), 1 deletion(-) diff --git a/src/mnemosyne/gateway.py b/src/mnemosyne/gateway.py index 070763c..2bd9dd2 100644 --- a/src/mnemosyne/gateway.py +++ b/src/mnemosyne/gateway.py @@ -92,6 +92,155 @@ def find_free_port() -> int: return s.getsockname()[1] +class YuyayStreamFilter: + """Strips yuyay protocol tags from SSE text_delta events. + + The cooperative memory protocol uses XML-like tags (, + , ) for gateway ↔ model sideband + communication. These must be stripped before the SSE stream reaches + the client (opencode) because they are internal to the proxy. + + Handles tags that span across multiple text_delta SSE events by + maintaining state across feed() calls. + """ + + _TAG_PAIRS: list[tuple[str, str]] = [ + ("", ""), + ("", ""), + ("", ""), + ] + + def __init__(self) -> None: + self._inside: str | None = None # closing tag we're waiting for + self._pending = "" # possible partial opening tag at end of chunk + self._sse_buf = b"" # incomplete SSE event bytes + + def feed(self, chunk: bytes) -> bytes: + """Process raw SSE bytes. Returns filtered bytes to forward.""" + self._sse_buf += chunk + out_parts: list[bytes] = [] + + # Process complete SSE events (delimited by \n\n) + while b"\n\n" in self._sse_buf: + event_bytes, self._sse_buf = self._sse_buf.split(b"\n\n", 1) + filtered = self._filter_event(event_bytes) + if filtered is not None: + out_parts.append(filtered + b"\n\n") + + return b"".join(out_parts) + + def flush(self) -> bytes: + """Flush any remaining buffered bytes (call at end of stream).""" + if self._sse_buf: + rest = self._sse_buf + self._sse_buf = b"" + return rest + return b"" + + def _filter_event(self, event_bytes: bytes) -> bytes | None: + """Filter a single complete SSE event. Returns None to suppress.""" + event_text = event_bytes.decode("utf-8", errors="replace") + data_str = None + prefix_lines: list[str] = [] + + for line in event_text.split("\n"): + if line.startswith("data: "): + data_str = line[6:] + else: + prefix_lines.append(line) + + if not data_str or data_str == "[DONE]": + return event_bytes + + try: + event_data = json.loads(data_str) + except json.JSONDecodeError: + return event_bytes + + # Only filter text_delta events + if ( + event_data.get("type") == "content_block_delta" + and event_data.get("delta", {}).get("type") == "text_delta" + ): + raw_text = event_data["delta"].get("text", "") + clean_text = self._filter_text(raw_text) + + if clean_text == raw_text: + return event_bytes # unchanged + + if not clean_text: + return None # entire chunk was tag content — suppress + + # Re-serialize with filtered text + event_data["delta"]["text"] = clean_text + new_data = json.dumps(event_data) + parts = prefix_lines + [f"data: {new_data}"] + return "\n".join(parts).encode("utf-8") + + return event_bytes + + def _filter_text(self, text: str) -> str: + """Strip yuyay tags from a text chunk, handling partial tags.""" + if self._pending: + text = self._pending + text + self._pending = "" + + result: list[str] = [] + i = 0 + + while i < len(text): + if self._inside: + # Inside a tag — look for closing tag + close_pos = text.find(self._inside, i) + if close_pos >= 0: + i = close_pos + len(self._inside) + self._inside = None + # Skip optional trailing newline + if i < len(text) and text[i] == "\n": + i += 1 + else: + # Still inside — consume remainder + break + else: + # Look for earliest opening tag + earliest_pos = -1 + matched_open = "" + matched_close = "" + for open_tag, close_tag in self._TAG_PAIRS: + pos = text.find(open_tag, i) + if pos >= 0 and (earliest_pos < 0 or pos < earliest_pos): + earliest_pos = pos + matched_open = open_tag + matched_close = close_tag + + if earliest_pos >= 0: + result.append(text[i:earliest_pos]) + self._inside = matched_close + i = earliest_pos + len(matched_open) + else: + # Check for partial opening tag at end of chunk + partial = self._find_partial(text, i) + if partial >= 0: + result.append(text[i:partial]) + self._pending = text[partial:] + break + else: + result.append(text[i:]) + break + + return "".join(result) + + def _find_partial(self, text: str, start: int) -> int: + """Find position of a partial opening tag at end of text.""" + for open_tag, _ in self._TAG_PAIRS: + for length in range(len(open_tag) - 1, 0, -1): + if text.endswith(open_tag[:length]): + pos = len(text) - length + if pos >= start: + return pos + return -1 + + def _session_id(body: dict) -> str: """Derive a stable session fingerprint from the conversation. @@ -1877,6 +2026,7 @@ def create_app( stream_error = False sse_buffer = bytearray() usage: dict[str, Any] = {} + yuyay_filter = YuyayStreamFilter() try: with client.stream( "POST", upstream_path, json=outgoing_body, headers=headers @@ -1890,6 +2040,7 @@ def create_app( for chunk in resp.iter_bytes(): bytes_out += len(chunk) chunk_count += 1 + # Inspect raw chunk for metrics (before filtering) _inspect_sse_chunk( chunk, buffer=sse_buffer, @@ -1899,7 +2050,14 @@ def create_app( provider=provider, usage_accumulator=usage, ) - yield chunk + # Filter yuyay protocol tags before forwarding + filtered = yuyay_filter.feed(chunk) + if filtered: + yield filtered + # Flush any remaining buffered bytes + tail = yuyay_filter.flush() + if tail: + yield tail except Exception as e: stream_error = True emit_event(