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]
|
||||
|
||||
|
||||
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(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue