websocket add
This commit is contained in:
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user