feat: add SSE stream filter for yuyay protocol tags

Strip <yuyay-response>, <yuyay-manifest>, and <yuyay-query> 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 <clio-agent@sisyphuslabs.ai>
This commit is contained in:
Joey Yakimowich-Payne 2026-03-13 11:49:12 -06:00
commit bee90915db

View file

@ -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 (<yuyay-response>,
<yuyay-manifest>, <yuyay-query>) 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]] = [
("<yuyay-response>", "</yuyay-response>"),
("<yuyay-manifest>", "</yuyay-manifest>"),
("<yuyay-query>", "</yuyay-query>"),
]
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(