"""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