Skip to content
File

Blob: src/worker/durable/run-do/websocket.ts

typescript190 lines
1import { serializeLogEvent, serializeRunExecutionState, serializeRunStep } from "@/worker/presentation/serializers";
2import { readInternalRunLogAuth } from "@/worker/run-logs/auth";
3import type { Logger } from "@/worker/services/logger";
4import { RunId, UserId } from "@/contracts";
5import type { RunWsLogMessage, RunWsStateMessage } from "@/contracts/run-ws";
6import type { RunLogRecord } from "@/worker/contracts";
7 
8import { listRunLogs } from "./logs";
9import type { RunDb } from "./repo";
10import { getRunMeta, listRunSteps } from "./repo/index";
11 
12export interface RunLogSocketAttachment {
13 runId: RunId;
14 userId: UserId;
15 connectedAt: number;
16 lastAckedSeq: number | null;
17}
18 
19export const getSocketAttachment = (ws: WebSocket): RunLogSocketAttachment | null => {
20 const attachment = ws.deserializeAttachment();
21 if (
22 !attachment ||
23 typeof attachment !== "object" ||
24 Array.isArray(attachment) ||
25 !("runId" in attachment) ||
26 !("userId" in attachment) ||
27 !("connectedAt" in attachment) ||
28 !("lastAckedSeq" in attachment)
29 ) {
30 return null;
31 }
32 
33 try {
34 return {
35 runId: RunId.assertDecode(attachment.runId),
36 userId: UserId.assertDecode(attachment.userId),
37 connectedAt: Number(attachment.connectedAt),
38 lastAckedSeq: attachment.lastAckedSeq === null ? null : Number(attachment.lastAckedSeq),
39 };
40 } catch {
41 return null;
42 }
43};
44 
45export const sendLogEvent = (ws: WebSocket, event: RunLogRecord): void => {
46 const message: RunWsLogMessage = { type: "log", event: serializeLogEvent(event) };
47 ws.send(JSON.stringify(message));
48};
49 
50const buildStateMessage = (
51 meta: NonNullable<Awaited<ReturnType<typeof getRunMeta>>>,
52 steps: Awaited<ReturnType<typeof listRunSteps>>,
53): RunWsStateMessage => ({
54 type: "state",
55 run: serializeRunExecutionState(meta),
56 steps: steps.map(serializeRunStep),
57});
58 
59const sendStateMessage = (ws: WebSocket, message: RunWsStateMessage): void => {
60 ws.send(JSON.stringify(message));
61};
62 
63export const broadcastStateUpdate = async (
64 ctx: DurableObjectState,
65 db: RunDb,
66 logger: Logger,
67 runId: RunId,
68): Promise<void> => {
69 const meta = await getRunMeta(db, runId);
70 if (!meta) return;
71 
72 const steps = await listRunSteps(db, runId);
73 const message = buildStateMessage(meta, steps);
74 
75 for (const ws of ctx.getWebSockets(runId)) {
76 const attachment = getSocketAttachment(ws);
77 if (!attachment || attachment.runId !== runId) continue;
78 
79 try {
80 sendStateMessage(ws, message);
81 } catch (error) {
82 logger.warn("run_state_socket_send_failed", {
83 runId,
84 userId: attachment.userId,
85 error: error instanceof Error ? error.message : String(error),
86 });
87 ws.close(1011, "state_delivery_failed");
88 }
89 }
90};
91 
92export const broadcastLogEvents = (
93 ctx: DurableObjectState,
94 logger: Logger,
95 runId: RunId,
96 events: readonly RunLogRecord[],
97): void => {
98 if (events.length === 0) {
99 return;
100 }
101 
102 for (const ws of ctx.getWebSockets(runId)) {
103 const attachment = getSocketAttachment(ws);
104 if (!attachment || attachment.runId !== runId) {
105 continue;
106 }
107 
108 try {
109 for (const event of events) {
110 sendLogEvent(ws, event);
111 }
112 } catch (error) {
113 logger.warn("run_log_socket_send_failed", {
114 runId,
115 userId: attachment.userId,
116 error: error instanceof Error ? error.message : String(error),
117 });
118 ws.close(1011, "log_delivery_failed");
119 }
120 }
121};
122 
123export const handleRunLogStreamFetch = async (
124 ctx: DurableObjectState,
125 db: RunDb,
126 request: Request,
127): Promise<Response> => {
128 const auth = readInternalRunLogAuth(request);
129 if (!auth) {
130 return new Response("Forbidden", { status: 403 });
131 }
132 
133 const url = new URL(request.url);
134 const match = /^\/api\/private\/runs\/([^/]+)\/logs$/u.exec(url.pathname);
135 if (request.method !== "GET" || !match) {
136 return new Response("Not found", { status: 404 });
137 }
138 
139 let runId: RunId;
140 try {
141 runId = RunId.assertDecode(match[1]);
142 } catch {
143 return new Response("Not found", { status: 404 });
144 }
145 
146 if (runId !== auth.runId) {
147 return new Response("Forbidden", { status: 403 });
148 }
149 
150 if (request.headers.get("upgrade")?.toLowerCase() !== "websocket") {
151 return new Response("WebSocket upgrade required", { status: 426 });
152 }
153 
154 const pair = new WebSocketPair();
155 const client = pair[0];
156 const server = pair[1];
157 
158 server.serializeAttachment({
159 runId,
160 userId: auth.userId,
161 connectedAt: Date.now(),
162 lastAckedSeq: null,
163 } satisfies RunLogSocketAttachment);
164 ctx.acceptWebSocket(server, [runId]);
165 
166 for (const event of await listRunLogs(db, runId)) {
167 sendLogEvent(server, event);
168 }
169 
170 const meta = await getRunMeta(db, runId);
171 if (meta) {
172 const steps = await listRunSteps(db, runId);
173 sendStateMessage(server, buildStateMessage(meta, steps));
174 }
175 
176 return new Response(null, {
177 status: 101,
178 webSocket: client,
179 });
180};
181 
182export const logRunSocketError = (logger: Logger, ws: WebSocket, error: unknown): void => {
183 const attachment = getSocketAttachment(ws);
184 logger.warn("run_log_socket_error", {
185 runId: attachment?.runId ?? null,
186 userId: attachment?.userId ?? null,
187 error: error instanceof Error ? error.message : String(error),
188 });
189};