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:
parent
6b9b3df64d
commit
bee90915db
1 changed files with 159 additions and 1 deletions
|
|
@ -92,6 +92,155 @@ def find_free_port() -> int:
|
||||||
return s.getsockname()[1]
|
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:
|
def _session_id(body: dict) -> str:
|
||||||
"""Derive a stable session fingerprint from the conversation.
|
"""Derive a stable session fingerprint from the conversation.
|
||||||
|
|
||||||
|
|
@ -1877,6 +2026,7 @@ def create_app(
|
||||||
stream_error = False
|
stream_error = False
|
||||||
sse_buffer = bytearray()
|
sse_buffer = bytearray()
|
||||||
usage: dict[str, Any] = {}
|
usage: dict[str, Any] = {}
|
||||||
|
yuyay_filter = YuyayStreamFilter()
|
||||||
try:
|
try:
|
||||||
with client.stream(
|
with client.stream(
|
||||||
"POST", upstream_path, json=outgoing_body, headers=headers
|
"POST", upstream_path, json=outgoing_body, headers=headers
|
||||||
|
|
@ -1890,6 +2040,7 @@ def create_app(
|
||||||
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
|
||||||
|
# Inspect raw chunk for metrics (before filtering)
|
||||||
_inspect_sse_chunk(
|
_inspect_sse_chunk(
|
||||||
chunk,
|
chunk,
|
||||||
buffer=sse_buffer,
|
buffer=sse_buffer,
|
||||||
|
|
@ -1899,7 +2050,14 @@ def create_app(
|
||||||
provider=provider,
|
provider=provider,
|
||||||
usage_accumulator=usage,
|
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:
|
except Exception as e:
|
||||||
stream_error = True
|
stream_error = True
|
||||||
emit_event(
|
emit_event(
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue