switch training progress from SSE to WebSocket
This commit is contained in:
parent
edf9d28930
commit
f4382e288b
3 changed files with 161 additions and 238 deletions
|
|
@ -27,11 +27,6 @@ logger = structlog.get_logger(__name__)
|
|||
|
||||
class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next: Callable) -> Response:
|
||||
# Bypass middleware for SSE streaming — BaseHTTPMiddleware
|
||||
# can buffer StreamingResponse bodies, breaking real-time delivery
|
||||
if request.url.path == "/api/train/progress":
|
||||
return await call_next(request)
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -7,8 +7,9 @@ Training API routes
|
|||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import StreamingResponse
|
||||
from starlette.websockets import WebSocketState
|
||||
from typing import Dict, Optional, Any
|
||||
import structlog
|
||||
from loggers import get_logger
|
||||
|
|
@ -494,101 +495,108 @@ async def get_training_metrics(
|
|||
)
|
||||
|
||||
|
||||
@router.get("/progress")
|
||||
async def stream_training_progress(
|
||||
request: Request,
|
||||
current_subject: str = Depends(get_current_subject),
|
||||
):
|
||||
@router.websocket("/progress/ws")
|
||||
async def ws_training_progress(websocket: WebSocket):
|
||||
"""
|
||||
Stream training progress updates using Server-Sent Events (SSE).
|
||||
Stream training progress updates over WebSocket.
|
||||
|
||||
This endpoint provides real-time updates on training progress.
|
||||
Supports reconnection via the SSE spec:
|
||||
- Sends `id:` with each event so the browser tracks position.
|
||||
- Sends `retry:` to control reconnection interval.
|
||||
- Sends named `event:` types (progress, heartbeat, complete, error).
|
||||
- Reads `Last-Event-ID` header on reconnect to replay missed steps.
|
||||
Supports reconnection via query params:
|
||||
- `last_event_id`: resume from a specific step on reconnect.
|
||||
- `token`: JWT auth token (WebSocket can't use Authorization header).
|
||||
"""
|
||||
# Read Last-Event-ID header for reconnection resume
|
||||
last_event_id = request.headers.get("last-event-id")
|
||||
# Auth: WebSocket can't use Authorization header, so accept token as query param
|
||||
from auth.authentication import is_auth_disabled
|
||||
|
||||
if not is_auth_disabled():
|
||||
token = websocket.query_params.get("token")
|
||||
if not token:
|
||||
await websocket.close(code = 4001, reason = "Missing auth token")
|
||||
return
|
||||
from auth.authentication import _decode_subject_without_verification
|
||||
from auth.storage import get_user_and_secret
|
||||
|
||||
subject = _decode_subject_without_verification(token)
|
||||
if subject is None:
|
||||
await websocket.close(code = 4001, reason = "Invalid token")
|
||||
return
|
||||
record = get_user_and_secret(subject)
|
||||
if record is None:
|
||||
await websocket.close(code = 4001, reason = "Invalid or expired token")
|
||||
return
|
||||
|
||||
await websocket.accept()
|
||||
|
||||
# Read optional last_event_id from query params for reconnection resume
|
||||
last_event_id = websocket.query_params.get("last_event_id")
|
||||
resume_from_step: Optional[int] = None
|
||||
if last_event_id is not None:
|
||||
try:
|
||||
resume_from_step = int(last_event_id)
|
||||
logger.info(f"SSE reconnect: resuming from step {resume_from_step}")
|
||||
logger.info(f"WebSocket reconnect: resuming from step {resume_from_step}")
|
||||
except ValueError:
|
||||
logger.warning(f"Invalid Last-Event-ID: {last_event_id}")
|
||||
logger.warning(f"Invalid last_event_id: {last_event_id}")
|
||||
|
||||
async def event_generator():
|
||||
backend = get_training_backend()
|
||||
job_id: str = getattr(backend, "current_job_id", "") or ""
|
||||
backend = get_training_backend()
|
||||
job_id: str = getattr(backend, "current_job_id", "") or ""
|
||||
|
||||
# ── Helpers ──────────────────────────────────────────────
|
||||
def build_progress(
|
||||
step: int,
|
||||
loss: float,
|
||||
learning_rate: float,
|
||||
total_steps: int,
|
||||
epoch: Optional[float] = None,
|
||||
progress: Optional[Any] = None,
|
||||
grad_norm_override: Optional[float] = None,
|
||||
eval_loss_override: Optional[float] = None,
|
||||
) -> TrainingProgress:
|
||||
total = max(total_steps, 0)
|
||||
if step < 0 or total == 0:
|
||||
progress_percent = 0.0
|
||||
else:
|
||||
progress_percent = (
|
||||
float(step) / float(total) * 100.0 if total > 0 else 0.0
|
||||
)
|
||||
|
||||
# Get actual values from progress object if available
|
||||
elapsed_seconds = (
|
||||
getattr(progress, "elapsed_seconds", None) if progress else None
|
||||
)
|
||||
eta_seconds = getattr(progress, "eta_seconds", None) if progress else None
|
||||
grad_norm = grad_norm_override
|
||||
if grad_norm is None and progress:
|
||||
grad_norm = getattr(progress, "grad_norm", None)
|
||||
num_tokens = getattr(progress, "num_tokens", None) if progress else None
|
||||
eval_loss = eval_loss_override
|
||||
if eval_loss is None and progress:
|
||||
eval_loss = getattr(progress, "eval_loss", None)
|
||||
|
||||
return TrainingProgress(
|
||||
job_id = job_id,
|
||||
step = step,
|
||||
total_steps = total,
|
||||
loss = loss,
|
||||
learning_rate = learning_rate,
|
||||
progress_percent = progress_percent,
|
||||
epoch = epoch,
|
||||
elapsed_seconds = elapsed_seconds,
|
||||
eta_seconds = eta_seconds,
|
||||
grad_norm = grad_norm,
|
||||
num_tokens = num_tokens,
|
||||
eval_loss = eval_loss,
|
||||
# ── Helpers ──────────────────────────────────────────────
|
||||
def build_progress(
|
||||
step: int,
|
||||
loss: float,
|
||||
learning_rate: float,
|
||||
total_steps: int,
|
||||
epoch: Optional[float] = None,
|
||||
progress: Optional[Any] = None,
|
||||
grad_norm_override: Optional[float] = None,
|
||||
eval_loss_override: Optional[float] = None,
|
||||
) -> TrainingProgress:
|
||||
total = max(total_steps, 0)
|
||||
if step < 0 or total == 0:
|
||||
progress_percent = 0.0
|
||||
else:
|
||||
progress_percent = (
|
||||
float(step) / float(total) * 100.0 if total > 0 else 0.0
|
||||
)
|
||||
|
||||
def format_sse(
|
||||
data: str,
|
||||
event: str = "progress",
|
||||
event_id: Optional[int] = None,
|
||||
) -> str:
|
||||
"""Format a single SSE message with id/event/data fields."""
|
||||
lines = []
|
||||
if event_id is not None:
|
||||
lines.append(f"id: {event_id}")
|
||||
lines.append(f"event: {event}")
|
||||
lines.append(f"data: {data}")
|
||||
lines.append("") # trailing blank line
|
||||
lines.append("") # double newline terminates the event
|
||||
return "\n".join(lines)
|
||||
elapsed_seconds = (
|
||||
getattr(progress, "elapsed_seconds", None) if progress else None
|
||||
)
|
||||
eta_seconds = getattr(progress, "eta_seconds", None) if progress else None
|
||||
grad_norm = grad_norm_override
|
||||
if grad_norm is None and progress:
|
||||
grad_norm = getattr(progress, "grad_norm", None)
|
||||
num_tokens = getattr(progress, "num_tokens", None) if progress else None
|
||||
eval_loss = eval_loss_override
|
||||
if eval_loss is None and progress:
|
||||
eval_loss = getattr(progress, "eval_loss", None)
|
||||
|
||||
# ── Retry directive ──────────────────────────────────────
|
||||
# Tell the browser to reconnect after 3 seconds if the connection drops
|
||||
yield "retry: 3000\n\n"
|
||||
return TrainingProgress(
|
||||
job_id = job_id,
|
||||
step = step,
|
||||
total_steps = total,
|
||||
loss = loss,
|
||||
learning_rate = learning_rate,
|
||||
progress_percent = progress_percent,
|
||||
epoch = epoch,
|
||||
elapsed_seconds = elapsed_seconds,
|
||||
eta_seconds = eta_seconds,
|
||||
grad_norm = grad_norm,
|
||||
num_tokens = num_tokens,
|
||||
eval_loss = eval_loss,
|
||||
)
|
||||
|
||||
async def send_event(event: str, event_id: Optional[int], payload: TrainingProgress):
|
||||
"""Send a typed event over WebSocket as JSON."""
|
||||
if websocket.client_state != WebSocketState.CONNECTED:
|
||||
return
|
||||
await websocket.send_json({
|
||||
"event": event,
|
||||
"id": event_id,
|
||||
"data": payload.model_dump(),
|
||||
})
|
||||
|
||||
try:
|
||||
# ── Replay missed steps on reconnect ─────────────────────
|
||||
if resume_from_step is not None and backend.step_history:
|
||||
replayed = 0
|
||||
|
|
@ -629,12 +637,10 @@ async def stream_training_progress(
|
|||
progress = tp_replay,
|
||||
grad_norm_override = grad_norm_by_step.get(step_val),
|
||||
)
|
||||
yield format_sse(
|
||||
payload.model_dump_json(), event = "progress", event_id = step_val
|
||||
)
|
||||
await send_event("progress", step_val, payload)
|
||||
replayed += 1
|
||||
if replayed:
|
||||
logger.info(f"SSE reconnect: replayed {replayed} missed steps")
|
||||
logger.info(f"WebSocket reconnect: replayed {replayed} missed steps")
|
||||
|
||||
# ── Initial status (only on fresh connections) ───────────
|
||||
if resume_from_step is None:
|
||||
|
|
@ -651,11 +657,9 @@ async def stream_training_progress(
|
|||
epoch = initial_epoch,
|
||||
progress = tp,
|
||||
)
|
||||
yield format_sse(
|
||||
initial_progress.model_dump_json(), event = "progress", event_id = 0
|
||||
)
|
||||
await send_event("progress", 0, initial_progress)
|
||||
|
||||
# If not active, send final state and exit
|
||||
# If not active, send final state and close
|
||||
if not is_active:
|
||||
if backend.step_history:
|
||||
final_step = backend.step_history[-1]
|
||||
|
|
@ -675,23 +679,18 @@ async def stream_training_progress(
|
|||
final_epoch,
|
||||
progress = tp,
|
||||
)
|
||||
yield format_sse(
|
||||
payload.model_dump_json(), event = "complete", event_id = final_step
|
||||
)
|
||||
await send_event("complete", final_step, payload)
|
||||
else:
|
||||
yield format_sse(
|
||||
build_progress(-1, 0.0, 0.0, 0, progress = tp).model_dump_json(),
|
||||
event = "complete",
|
||||
event_id = 0,
|
||||
await send_event(
|
||||
"complete", 0,
|
||||
build_progress(-1, 0.0, 0.0, 0, progress = tp),
|
||||
)
|
||||
return
|
||||
|
||||
# ── Live polling loop ────────────────────────────────────
|
||||
last_step = resume_from_step if resume_from_step is not None else -1
|
||||
no_update_count = 0
|
||||
max_no_updates = (
|
||||
1800 # Timeout after 30 minutes (large models need time for compilation)
|
||||
)
|
||||
max_no_updates = 1800 # Timeout after 30 min
|
||||
|
||||
while backend.is_training_active():
|
||||
try:
|
||||
|
|
@ -723,11 +722,7 @@ async def stream_training_progress(
|
|||
current_epoch,
|
||||
progress = tp_inner,
|
||||
)
|
||||
yield format_sse(
|
||||
progress_payload.model_dump_json(),
|
||||
event = "progress",
|
||||
event_id = current_step,
|
||||
)
|
||||
await send_event("progress", current_step, progress_payload)
|
||||
last_step = current_step
|
||||
no_update_count = 0
|
||||
else:
|
||||
|
|
@ -742,17 +737,11 @@ async def stream_training_progress(
|
|||
current_epoch,
|
||||
progress = tp_inner,
|
||||
)
|
||||
yield format_sse(
|
||||
heartbeat_payload.model_dump_json(),
|
||||
event = "heartbeat",
|
||||
event_id = current_step,
|
||||
)
|
||||
await send_event("heartbeat", current_step, heartbeat_payload)
|
||||
else:
|
||||
# No steps yet, but training is active (model loading, etc.)
|
||||
no_update_count += 1
|
||||
if no_update_count % 5 == 0:
|
||||
# Pull total_steps and status from trainer so
|
||||
# the frontend can show "Tokenizing…" etc.
|
||||
tp_prep = getattr(
|
||||
getattr(backend, "trainer", None),
|
||||
"training_progress",
|
||||
|
|
@ -762,17 +751,9 @@ async def stream_training_progress(
|
|||
getattr(tp_prep, "total_steps", 0) if tp_prep else 0
|
||||
)
|
||||
preparing_payload = build_progress(
|
||||
0,
|
||||
0.0,
|
||||
0.0,
|
||||
prep_total,
|
||||
progress = tp_prep,
|
||||
)
|
||||
yield format_sse(
|
||||
preparing_payload.model_dump_json(),
|
||||
event = "heartbeat",
|
||||
event_id = 0,
|
||||
0, 0.0, 0.0, prep_total, progress = tp_prep,
|
||||
)
|
||||
await send_event("heartbeat", 0, preparing_payload)
|
||||
|
||||
# Timeout check
|
||||
if no_update_count > max_no_updates:
|
||||
|
|
@ -783,10 +764,8 @@ async def stream_training_progress(
|
|||
timeout_payload = build_progress(
|
||||
last_step, 0.0, 0.0, 0, progress = tp_timeout
|
||||
)
|
||||
yield format_sse(
|
||||
timeout_payload.model_dump_json(),
|
||||
event = "error",
|
||||
event_id = last_step if last_step >= 0 else 0,
|
||||
await send_event(
|
||||
"error", last_step if last_step >= 0 else 0, timeout_payload
|
||||
)
|
||||
break
|
||||
|
||||
|
|
@ -798,10 +777,8 @@ async def stream_training_progress(
|
|||
getattr(backend, "trainer", None), "training_progress", None
|
||||
)
|
||||
error_payload = build_progress(0, 0.0, 0.0, 0, progress = tp_error)
|
||||
yield format_sse(
|
||||
error_payload.model_dump_json(),
|
||||
event = "error",
|
||||
event_id = last_step if last_step >= 0 else 0,
|
||||
await send_event(
|
||||
"error", last_step if last_step >= 0 else 0, error_payload
|
||||
)
|
||||
break
|
||||
|
||||
|
|
@ -822,18 +799,16 @@ async def stream_training_progress(
|
|||
final_epoch,
|
||||
progress = final_tp,
|
||||
)
|
||||
yield format_sse(
|
||||
final_payload.model_dump_json(),
|
||||
event = "complete",
|
||||
event_id = final_step if final_step >= 0 else 0,
|
||||
await send_event(
|
||||
"complete", final_step if final_step >= 0 else 0, final_payload
|
||||
)
|
||||
|
||||
return StreamingResponse(
|
||||
event_generator(),
|
||||
media_type = "text/event-stream",
|
||||
headers = {
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
"X-Accel-Buffering": "no",
|
||||
},
|
||||
)
|
||||
except WebSocketDisconnect:
|
||||
logger.info("WebSocket client disconnected")
|
||||
except Exception as e:
|
||||
logger.error(f"WebSocket error: {e}", exc_info = True)
|
||||
try:
|
||||
if websocket.client_state == WebSocketState.CONNECTED:
|
||||
await websocket.send_json({"event": "error", "id": None, "data": {"error": str(e)}})
|
||||
except Exception:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -78,110 +78,63 @@ interface ParsedSseEvent {
|
|||
id: number | null;
|
||||
}
|
||||
|
||||
function parseSseEvent(rawEvent: string): ParsedSseEvent | null {
|
||||
const lines = rawEvent.split(/\r?\n/);
|
||||
let eventName: ProgressEventName = "progress";
|
||||
let id: number | null = null;
|
||||
const dataLines: string[] = [];
|
||||
|
||||
for (const line of lines) {
|
||||
if (!line) {
|
||||
continue;
|
||||
}
|
||||
if (line.startsWith("event:")) {
|
||||
const value = line.slice(6).trim();
|
||||
if (
|
||||
value === "progress" ||
|
||||
value === "heartbeat" ||
|
||||
value === "complete" ||
|
||||
value === "error"
|
||||
) {
|
||||
eventName = value;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if (line.startsWith("id:")) {
|
||||
const value = Number(line.slice(3).trim());
|
||||
id = Number.isFinite(value) ? value : null;
|
||||
continue;
|
||||
}
|
||||
if (line.startsWith("data:")) {
|
||||
dataLines.push(line.slice(5).trimStart());
|
||||
}
|
||||
}
|
||||
|
||||
if (dataLines.length === 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const parsed = JSON.parse(dataLines.join("\n")) as TrainingProgressPayload;
|
||||
return { event: eventName, payload: parsed, id };
|
||||
}
|
||||
|
||||
export async function streamTrainingProgress(options: {
|
||||
signal: AbortSignal;
|
||||
lastEventId?: number | null;
|
||||
onOpen?: () => void;
|
||||
onEvent: (event: ParsedSseEvent) => void;
|
||||
}): Promise<void> {
|
||||
const headers = new Headers();
|
||||
// Build WebSocket URL from current page location
|
||||
const protocol = window.location.protocol === "https:" ? "wss:" : "ws:";
|
||||
const params = new URLSearchParams();
|
||||
|
||||
// Pass auth token as query param (WebSocket can't use Authorization header)
|
||||
const token = localStorage.getItem("unsloth_auth_token");
|
||||
if (token) params.set("token", token);
|
||||
if (typeof options.lastEventId === "number") {
|
||||
headers.set("Last-Event-ID", String(options.lastEventId));
|
||||
params.set("last_event_id", String(options.lastEventId));
|
||||
}
|
||||
|
||||
const response = await authFetch("/api/train/progress", {
|
||||
method: "GET",
|
||||
headers,
|
||||
signal: options.signal,
|
||||
});
|
||||
const url = `${protocol}//${window.location.host}/api/train/progress/ws?${params}`;
|
||||
|
||||
if (!response.ok) {
|
||||
throw new Error(await readError(response));
|
||||
}
|
||||
return new Promise<void>((resolve, reject) => {
|
||||
const ws = new WebSocket(url);
|
||||
|
||||
if (!response.body) {
|
||||
throw new Error("Progress stream unavailable");
|
||||
}
|
||||
// Wire up AbortSignal to close the socket
|
||||
const onAbort = () => ws.close();
|
||||
options.signal.addEventListener("abort", onAbort);
|
||||
|
||||
options.onOpen?.();
|
||||
|
||||
const reader = response.body.getReader();
|
||||
const decoder = new TextDecoder();
|
||||
let buffer = "";
|
||||
|
||||
while (true) {
|
||||
const { value, done } = await reader.read();
|
||||
if (done) {
|
||||
break;
|
||||
}
|
||||
|
||||
buffer += decoder.decode(value, { stream: true });
|
||||
|
||||
let separatorIndex = buffer.search(/\r?\n\r?\n/);
|
||||
while (separatorIndex >= 0) {
|
||||
const rawEvent = buffer.slice(0, separatorIndex);
|
||||
const separatorLength = buffer[separatorIndex] === "\r" ? 4 : 2;
|
||||
buffer = buffer.slice(separatorIndex + separatorLength);
|
||||
|
||||
if (rawEvent.startsWith("retry:")) {
|
||||
separatorIndex = buffer.search(/\r?\n\r?\n/);
|
||||
continue;
|
||||
}
|
||||
ws.onopen = () => {
|
||||
options.onOpen?.();
|
||||
};
|
||||
|
||||
ws.onmessage = (messageEvent) => {
|
||||
try {
|
||||
const event = parseSseEvent(rawEvent);
|
||||
if (event) {
|
||||
options.onEvent(event);
|
||||
}
|
||||
} catch (error) {
|
||||
if (!isAbortError(error)) {
|
||||
throw error;
|
||||
}
|
||||
const msg = JSON.parse(messageEvent.data) as {
|
||||
event: ProgressEventName;
|
||||
id: number | null;
|
||||
data: TrainingProgressPayload;
|
||||
};
|
||||
options.onEvent({
|
||||
event: msg.event,
|
||||
id: msg.id,
|
||||
payload: msg.data,
|
||||
});
|
||||
} catch {
|
||||
// Ignore parse errors for malformed messages
|
||||
}
|
||||
};
|
||||
|
||||
separatorIndex = buffer.search(/\r?\n\r?\n/);
|
||||
}
|
||||
}
|
||||
ws.onclose = () => {
|
||||
options.signal.removeEventListener("abort", onAbort);
|
||||
resolve();
|
||||
};
|
||||
|
||||
ws.onerror = () => {
|
||||
options.signal.removeEventListener("abort", onAbort);
|
||||
reject(new Error("WebSocket connection failed"));
|
||||
};
|
||||
});
|
||||
}
|
||||
|
||||
export { isAbortError };
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue