websocket add

This commit is contained in:
2026-05-31 17:17:30 +03:00
parent 53c0bd1342
commit f554aa1b5e
14 changed files with 590 additions and 141 deletions

View File

@@ -470,7 +470,10 @@ async def websocket_endpoint(
websocket: WebSocket, task_id: str, source: str = None, level: str = None
):
"""
WebSocket endpoint for real-time log streaming with optional server-side filtering.
WebSocket endpoint for real-time log streaming AND task status updates.
Sends two message types:
- Log entries: plain dicts with level/message/timestamp (backward compatible, no type field)
- Status updates: dict with type="task_status" and nested task dict
Query Parameters:
source: Filter logs by source component (e.g., "plugin", "superset_api")
level: Filter logs by minimum level (DEBUG, INFO, WARNING, ERROR)
@@ -490,7 +493,7 @@ async def websocket_endpoint(
level_hierarchy = {"DEBUG": 0, "INFO": 1, "WARNING": 2, "ERROR": 3}
min_level = level_hierarchy.get(level_filter, 0) if level_filter else 0
logger.reason(
"Accepted WebSocket log stream connection",
"Accepted WebSocket log+status stream connection",
extra={
"task_id": task_id,
"source_filter": source_filter,
@@ -499,11 +502,13 @@ async def websocket_endpoint(
},
)
task_manager = get_task_manager()
queue = await task_manager.subscribe_logs(task_id)
log_queue = await task_manager.subscribe_logs(task_id)
status_queue = await task_manager.subscribe_status(task_id)
logger.reason(
"Subscribed WebSocket client to task log queue",
"Subscribed WebSocket client to task log and status queues",
extra={"task_id": task_id},
)
def matches_filters(log_entry) -> bool:
"""Check if log entry matches the filter criteria."""
log_source = getattr(log_entry, "source", None)
@@ -514,7 +519,37 @@ async def websocket_endpoint(
if log_level < min_level:
return False
return True
async def send_status_update(task: Any) -> None:
"""Send a structured task status update over the WebSocket."""
status_dict = {
"id": task.id,
"plugin_id": task.plugin_id,
"status": task.status.value if hasattr(task.status, "value") else str(task.status),
"started_at": task.started_at.isoformat() if task.started_at else None,
"finished_at": task.finished_at.isoformat() if task.finished_at else None,
"user_id": task.user_id,
"result": task.result,
"input_required": task.input_required,
"input_request": task.input_request,
}
await websocket.send_json({
"type": "task_status",
"task_id": task.id,
"task": status_dict,
})
try:
# ── Send initial task status ──
task = task_manager.get_task(task_id)
if task:
await send_status_update(task)
logger.reason(
"Sent initial task status",
extra={"task_id": task_id, "status": str(task.status)},
)
# ── Replay initial logs ──
logger.reason(
"Starting task log stream replay and live forwarding",
extra={"task_id": task_id},
@@ -535,6 +570,8 @@ async def websocket_endpoint(
"total_available_logs": len(initial_logs),
},
)
# ── Send synthetic AWAITING_INPUT prompt if needed ──
task = task_manager.get_task(task_id)
if task and task.status == "AWAITING_INPUT" and task.input_request:
synthetic_log = {
@@ -550,47 +587,167 @@ async def websocket_endpoint(
"Replayed awaiting-input prompt to restored WebSocket client",
extra={"task_id": task_id, "task_status": task.status},
)
# ── Main loop: listen on both log and status queues ──
while True:
log_entry = await queue.get()
if not matches_filters(log_entry):
continue
log_dict = log_entry.model_dump()
log_dict["timestamp"] = log_dict["timestamp"].isoformat()
await websocket.send_json(log_dict)
logger.reflect(
"Forwarded task log entry to WebSocket client",
extra={
"task_id": task_id,
"level": log_dict.get("level"),
},
done, _ = await asyncio.wait(
[log_queue.get(), status_queue.get()],
return_when=asyncio.FIRST_COMPLETED,
)
if (
"Task completed successfully" in log_entry.message
or "Task failed" in log_entry.message
):
logger.reason(
"Observed terminal task log entry; delaying to preserve client visibility",
extra={"task_id": task_id, "message": log_entry.message},
for coro in done:
result = coro.result()
# ── Status update ──
if isinstance(result, dict) and result.get("type") == "task_status":
await websocket.send_json(result)
task_status = result.get("task", {}).get("status", "")
if task_status in ("SUCCESS", "FAILED"):
logger.reason(
"Task reached terminal state via status broadcast; closing stream",
extra={"task_id": task_id, "status": task_status},
)
await asyncio.sleep(2)
raise StopIteration # exit the while loop
continue
# ── Log entry ──
if not matches_filters(result):
continue
log_dict = result.model_dump()
log_dict["timestamp"] = log_dict["timestamp"].isoformat()
await websocket.send_json(log_dict)
logger.reflect(
"Forwarded task log entry to WebSocket client",
extra={
"task_id": task_id,
"level": log_dict.get("level"),
},
)
await asyncio.sleep(2)
except WebSocketDisconnect:
if (
"Task completed successfully" in result.message
or "Task failed" in result.message
):
logger.reason(
"Observed terminal task log entry; delaying to preserve client visibility",
extra={"task_id": task_id, "message": result.message},
)
await asyncio.sleep(2)
except (WebSocketDisconnect, StopIteration):
if isinstance(StopIteration):
pass # normal termination
logger.reason(
"WebSocket client disconnected from task log stream",
"WebSocket client disconnected or stream ended",
extra={"task_id": task_id},
)
except Exception as exc:
logger.explore(
"WebSocket log streaming encountered an unexpected failure",
"WebSocket log+status streaming encountered an unexpected failure",
extra={"task_id": task_id, "error": str(exc)},
)
raise
finally:
task_manager.unsubscribe_logs(task_id, queue)
task_manager.unsubscribe_logs(task_id, log_queue)
task_manager.unsubscribe_status(task_id, status_queue)
logger.reflect(
"Released WebSocket log queue subscription",
"Released WebSocket log and status queue subscriptions",
extra={"task_id": task_id},
)
# #endregion websocket_endpoint
# #region task_events_websocket [C:4] [TYPE Function]
# @BRIEF WebSocket endpoint for global task events (status changes for ALL tasks).
# @RELATION CALLS -> [TaskManagerPackage]
# @PRE WebSocket must be authenticated via `token` query param.
# @POST WebSocket streams task status events until disconnect.
@app.websocket("/ws/task-events")
async def task_events_websocket(websocket: WebSocket):
"""
WebSocket endpoint for global task events.
Streams {type: "task_status", task_id: ..., task: {...}} for ALL task status changes.
Query Parameters:
token: JWT or API key for authentication (required)
"""
seed_trace_id()
with belief_scope("task_events_websocket"):
if not await _authenticate_websocket(websocket, "ws/task-events"):
await websocket.close(code=4001, reason="Authentication required")
return
await websocket.accept()
logger.reason("Accepted global task events WebSocket connection")
task_manager = get_task_manager()
event_queue = await task_manager.subscribe_task_events()
logger.reason("Subscribed to global task events")
try:
while True:
event = await event_queue.get()
await websocket.send_json(event)
logger.reflect(
"Forwarded task event to global client",
extra={"task_id": event.get("task_id")},
)
except WebSocketDisconnect:
logger.reason("Global task events client disconnected")
except Exception as exc:
logger.explore(
"Global task events streaming failed",
extra={"error": str(exc)},
)
raise
finally:
task_manager.unsubscribe_task_events(event_queue)
logger.reflect("Released global task events subscription")
# #endregion task_events_websocket
# #region maintenance_events_websocket [C:4] [TYPE Function]
# @BRIEF WebSocket endpoint for maintenance events (created/ended/banner changes).
# @RELATION CALLS -> [TaskManagerPackage]
# @PRE WebSocket must be authenticated via `token` query param.
# @POST WebSocket streams maintenance events until disconnect.
@app.websocket("/ws/maintenance/events")
async def maintenance_events_websocket(websocket: WebSocket):
"""
WebSocket endpoint for maintenance events.
Streams {type: "maintenance.event_created", ...} / {type: "maintenance.event_ended"}.
Query Parameters:
token: JWT or API key for authentication (required)
"""
seed_trace_id()
with belief_scope("maintenance_events_websocket"):
if not await _authenticate_websocket(websocket, "ws/maintenance/events"):
await websocket.close(code=4001, reason="Authentication required")
return
await websocket.accept()
logger.reason("Accepted maintenance events WebSocket connection")
task_manager = get_task_manager()
event_queue = await task_manager.subscribe_maintenance_events()
logger.reason("Subscribed to maintenance events")
try:
while True:
event = await event_queue.get()
await websocket.send_json(event)
logger.reflect(
"Forwarded maintenance event to client",
extra={"event_type": event.get("type"), "maintenance_id": event.get("maintenance_id")},
)
except WebSocketDisconnect:
logger.reason("Maintenance events client disconnected")
except Exception as exc:
logger.explore(
"Maintenance events streaming failed",
extra={"error": str(exc)},
)
raise
finally:
task_manager.unsubscribe_maintenance_events(event_queue)
logger.reflect("Released maintenance events subscription")
# #endregion maintenance_events_websocket
# #region dataset_websocket_endpoint [C:4] [TYPE Function]
# @BRIEF WebSocket endpoint for dataset.updated events — auto-refresh on task completion.
# @RELATION CALLS -> [TaskManagerPackage]