feat: prepare ARR for controlled public deployment
This commit is contained in:
217
agent_integration/events.py
Normal file
217
agent_integration/events.py
Normal file
@@ -0,0 +1,217 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user