218 lines
7.6 KiB
Python
218 lines
7.6 KiB
Python
"""Normalize raw Open Agent SSE events into a small public event contract."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Any, Dict, List, Optional, Sequence
|
|
|
|
from agent_integration.client import OpenAgentEvent
|
|
|
|
|
|
PUBLIC_ERROR_KEYS = ("code", "message", "retryable", "run_id", "request_id")
|
|
|
|
|
|
class AgentStreamNormalizer:
|
|
"""Statefully filter LangGraph internals and expose only user-facing events."""
|
|
|
|
def __init__(self) -> None:
|
|
self.run_id: Optional[str] = None
|
|
self.final_content: Optional[str] = None
|
|
self.fatal_error = False
|
|
self._text_parts: List[str] = []
|
|
self._final_candidate: Optional[str] = None
|
|
self._pending_error: Optional[Dict[str, Any]] = None
|
|
self._run_status: Optional[str] = None
|
|
self._started = False
|
|
self._finished = False
|
|
|
|
def feed(self, event: OpenAgentEvent) -> List[OpenAgentEvent]:
|
|
if self._finished:
|
|
return []
|
|
|
|
event_name = event.event or ""
|
|
if event_name == "metadata":
|
|
return self._metadata_event(event)
|
|
if event_name == "message.delta":
|
|
content = _content_text(event.data)
|
|
return self._delta_event(content, event.event_id)
|
|
if event_name == "messages":
|
|
return self._messages_event(event)
|
|
if event_name == "values":
|
|
self._capture_values_candidate(event.data)
|
|
return []
|
|
if event_name == "error":
|
|
self._pending_error = _sanitize_error(event.data)
|
|
self._capture_run_id(self._pending_error)
|
|
return []
|
|
if event_name in {"run.completed", "run.failed", "run.cancelled"}:
|
|
if isinstance(event.data, dict):
|
|
status = event.data.get("status")
|
|
if isinstance(status, str) and status:
|
|
self._run_status = status
|
|
self._capture_run_id(event.data)
|
|
return []
|
|
if event_name == "end":
|
|
return self.finish(event_id=event.event_id)
|
|
return []
|
|
|
|
def finish(self, *, event_id: Optional[str] = None) -> List[OpenAgentEvent]:
|
|
if self._finished:
|
|
return []
|
|
self._finished = True
|
|
|
|
streamed_content = "".join(self._text_parts)
|
|
final_content = self._final_candidate or streamed_content or None
|
|
self.final_content = final_content
|
|
public_events: List[OpenAgentEvent] = []
|
|
|
|
if final_content:
|
|
public_events.append(
|
|
OpenAgentEvent(
|
|
event="message.completed",
|
|
data={
|
|
"content": final_content,
|
|
"streamed": bool(streamed_content),
|
|
},
|
|
event_id=event_id,
|
|
)
|
|
)
|
|
|
|
if self._pending_error is not None:
|
|
recovered = (
|
|
self._pending_error.get("code") == "open_agent_final_content_missing"
|
|
and bool(final_content)
|
|
)
|
|
public_events.append(
|
|
OpenAgentEvent(
|
|
event="run.warning" if recovered else "run.error",
|
|
data=self._pending_error,
|
|
event_id=event_id,
|
|
)
|
|
)
|
|
self.fatal_error = not recovered
|
|
|
|
if self.fatal_error:
|
|
status = "failed"
|
|
elif self._pending_error is not None:
|
|
status = "completed_with_warning"
|
|
else:
|
|
status = self._run_status or "completed"
|
|
end_data: Dict[str, Any] = {"status": status}
|
|
if self.run_id:
|
|
end_data["run_id"] = self.run_id
|
|
public_events.append(OpenAgentEvent(event="run.end", data=end_data, event_id=event_id))
|
|
return public_events
|
|
|
|
def _metadata_event(self, event: OpenAgentEvent) -> List[OpenAgentEvent]:
|
|
if not isinstance(event.data, dict):
|
|
return []
|
|
self._capture_run_id(event.data)
|
|
if self._started:
|
|
return []
|
|
self._started = True
|
|
data: Dict[str, Any] = {}
|
|
if self.run_id:
|
|
data["run_id"] = self.run_id
|
|
return [OpenAgentEvent(event="run.started", data=data, event_id=event.event_id)]
|
|
|
|
def _messages_event(self, event: OpenAgentEvent) -> List[OpenAgentEvent]:
|
|
if not isinstance(event.data, Sequence) or isinstance(event.data, (str, bytes)):
|
|
return []
|
|
if not event.data or not isinstance(event.data[0], dict):
|
|
return []
|
|
|
|
message = event.data[0]
|
|
trace = event.data[1] if len(event.data) > 1 and isinstance(event.data[1], dict) else {}
|
|
if _is_hidden_middleware(trace):
|
|
return []
|
|
|
|
message_type = str(message.get("type", "")).lower()
|
|
if message_type not in {"aimessagechunk", "ai", "assistant"}:
|
|
return []
|
|
content = _content_text(message.get("content"))
|
|
if not content:
|
|
return []
|
|
if message_type == "aimessagechunk":
|
|
return self._delta_event(content, event.event_id)
|
|
|
|
self._final_candidate = content
|
|
return []
|
|
|
|
def _delta_event(self, content: str, event_id: Optional[str]) -> List[OpenAgentEvent]:
|
|
if not content:
|
|
return []
|
|
self._text_parts.append(content)
|
|
return [
|
|
OpenAgentEvent(
|
|
event="message.delta",
|
|
data={"content": content},
|
|
event_id=event_id,
|
|
)
|
|
]
|
|
|
|
def _capture_values_candidate(self, data: Any) -> None:
|
|
if not isinstance(data, dict):
|
|
return
|
|
messages = data.get("messages")
|
|
if not isinstance(messages, list):
|
|
return
|
|
for message in reversed(messages):
|
|
if not isinstance(message, dict):
|
|
continue
|
|
message_type = str(message.get("type", "")).lower()
|
|
if message_type not in {"ai", "assistant"}:
|
|
continue
|
|
content = _content_text(message.get("content"))
|
|
if content:
|
|
self._final_candidate = content
|
|
return
|
|
|
|
def _capture_run_id(self, data: Dict[str, Any]) -> None:
|
|
run_id = data.get("run_id")
|
|
if isinstance(run_id, str) and run_id:
|
|
self.run_id = run_id
|
|
|
|
|
|
def _content_text(value: Any) -> str:
|
|
if isinstance(value, str):
|
|
return value
|
|
if isinstance(value, dict):
|
|
for key in ("content", "text", "delta"):
|
|
if key not in value:
|
|
continue
|
|
content = _content_text(value[key])
|
|
if content:
|
|
return content
|
|
return ""
|
|
if not isinstance(value, list):
|
|
return ""
|
|
text_parts: List[str] = []
|
|
for block in value:
|
|
if isinstance(block, str):
|
|
text_parts.append(block)
|
|
elif isinstance(block, dict):
|
|
text = block.get("text")
|
|
if isinstance(text, str):
|
|
text_parts.append(text)
|
|
return "".join(text_parts)
|
|
|
|
|
|
def _is_hidden_middleware(trace: Dict[str, Any]) -> bool:
|
|
node = str(trace.get("langgraph_node", ""))
|
|
if "TitleMiddleware" in node:
|
|
return True
|
|
tags = trace.get("tags")
|
|
if isinstance(tags, list):
|
|
return any(str(tag).startswith("middleware:title") for tag in tags)
|
|
return False
|
|
|
|
|
|
def _sanitize_error(data: Any) -> Dict[str, Any]:
|
|
if not isinstance(data, dict):
|
|
return {"code": "open_agent_stream_error", "message": str(data)}
|
|
sanitized = {key: data[key] for key in PUBLIC_ERROR_KEYS if key in data}
|
|
if "code" not in sanitized:
|
|
sanitized["code"] = "open_agent_stream_error"
|
|
if "message" not in sanitized:
|
|
sanitized["message"] = "Open Agent stream failed"
|
|
return sanitized
|