@rivetkit/engine-runner 0.0.0-0-0-0-preview-guard-stops.9d82529

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
package/src/tunnel.ts ADDED
@@ -0,0 +1,1170 @@
1
+ import type * as protocol from "@rivetkit/engine-runner-protocol";
2
+ import type { GatewayId, RequestId } from "@rivetkit/engine-runner-protocol";
3
+ import type { Logger } from "pino";
4
+ import { type Runner, type RunnerActor, RunnerShutdownError } from "./mod";
5
+ import {
6
+ stringifyToClientTunnelMessageKind,
7
+ stringifyToServerTunnelMessageKind,
8
+ } from "./stringify";
9
+ import {
10
+ arraysEqual,
11
+ idToStr,
12
+ MAX_PAYLOAD_SIZE,
13
+ stringifyError,
14
+ unreachable,
15
+ } from "./utils";
16
+ import {
17
+ HIBERNATABLE_SYMBOL,
18
+ WebSocketTunnelAdapter,
19
+ } from "./websocket-tunnel-adapter";
20
+
21
+ export interface PendingRequest {
22
+ resolve: (response: Response) => void;
23
+ reject: (error: Error) => void;
24
+ streamController?: ReadableStreamDefaultController<Uint8Array>;
25
+ actorId?: string;
26
+ gatewayId?: GatewayId;
27
+ requestId?: RequestId;
28
+ clientMessageIndex: number;
29
+ }
30
+
31
+ export interface HibernatingWebSocketMetadata {
32
+ gatewayId: GatewayId;
33
+ requestId: RequestId;
34
+ clientMessageIndex: number;
35
+ serverMessageIndex: number;
36
+
37
+ path: string;
38
+ headers: Record<string, string>;
39
+ }
40
+
41
+ export class Tunnel {
42
+ #runner: Runner;
43
+
44
+ /** Maps request IDs to actor IDs for lookup */
45
+ #requestToActor: Array<{
46
+ gatewayId: GatewayId;
47
+ requestId: RequestId;
48
+ actorId: string;
49
+ }> = [];
50
+
51
+ /** Buffer for messages when not connected */
52
+ #bufferedMessages: Array<{
53
+ gatewayId: GatewayId;
54
+ requestId: RequestId;
55
+ messageKind: protocol.ToServerTunnelMessageKind;
56
+ }> = [];
57
+
58
+ get log(): Logger | undefined {
59
+ return this.#runner.log;
60
+ }
61
+
62
+ constructor(runner: Runner) {
63
+ this.#runner = runner;
64
+ }
65
+
66
+ start(): void {
67
+ // No-op - kept for compatibility
68
+ }
69
+
70
+ resendBufferedEvents(): void {
71
+ if (this.#bufferedMessages.length === 0) {
72
+ return;
73
+ }
74
+
75
+ this.log?.info({
76
+ msg: "resending buffered tunnel messages",
77
+ count: this.#bufferedMessages.length,
78
+ });
79
+
80
+ const messages = this.#bufferedMessages;
81
+ this.#bufferedMessages = [];
82
+
83
+ for (const { gatewayId, requestId, messageKind } of messages) {
84
+ this.#sendMessage(gatewayId, requestId, messageKind);
85
+ }
86
+ }
87
+
88
+ shutdown() {
89
+ // NOTE: Pegboard WS already closed at this point, cannot send
90
+ // anything. All teardown logic is handled by pegboard-runner.
91
+
92
+ // Reject all pending requests and close all WebSockets for all actors
93
+ // RunnerShutdownError will be explicitly ignored
94
+ for (const [_actorId, actor] of this.#runner.actors) {
95
+ // Reject all pending requests for this actor
96
+ for (const entry of actor.pendingRequests) {
97
+ entry.request.reject(new RunnerShutdownError());
98
+ }
99
+ actor.pendingRequests = [];
100
+
101
+ // Close all WebSockets for this actor
102
+ // The WebSocket close event with retry is automatically sent when the
103
+ // runner WS closes, so we only need to notify the client that the WS
104
+ // closed:
105
+ // https://github.com/rivet-dev/rivet/blob/00d4f6a22da178a6f8115e5db50d96c6f8387c2e/engine/packages/pegboard-runner/src/lib.rs#L157
106
+ for (const entry of actor.webSockets) {
107
+ // Only close non-hibernatable websockets to prevent sending
108
+ // unnecessary close messages for websockets that will be hibernated
109
+ if (!entry.ws[HIBERNATABLE_SYMBOL]) {
110
+ entry.ws._closeWithoutCallback(1000, "ws.tunnel_shutdown");
111
+ }
112
+ }
113
+ actor.webSockets = [];
114
+ }
115
+
116
+ // Clear the request-to-actor mapping
117
+ this.#requestToActor = [];
118
+ }
119
+
120
+ async restoreHibernatingRequests(
121
+ actorId: string,
122
+ metaEntries: HibernatingWebSocketMetadata[],
123
+ ) {
124
+ const actor = this.#runner.getActor(actorId);
125
+ if (!actor) {
126
+ throw new Error(
127
+ `Actor ${actorId} not found for restoring hibernating requests`,
128
+ );
129
+ }
130
+
131
+ if (actor.hibernationRestored) {
132
+ throw new Error(
133
+ `Actor ${actorId} already restored hibernating requests`,
134
+ );
135
+ }
136
+
137
+ this.log?.debug({
138
+ msg: "restoring hibernating requests",
139
+ actorId,
140
+ requests: actor.hibernatingRequests.length,
141
+ });
142
+
143
+ // Track all background operations
144
+ const backgroundOperations: Promise<void>[] = [];
145
+
146
+ // Process connected WebSockets
147
+ let connectedButNotLoadedCount = 0;
148
+ let restoredCount = 0;
149
+ for (const { gatewayId, requestId } of actor.hibernatingRequests) {
150
+ const requestIdStr = idToStr(requestId);
151
+ const meta = metaEntries.find(
152
+ (entry) =>
153
+ arraysEqual(entry.gatewayId, gatewayId) &&
154
+ arraysEqual(entry.requestId, requestId),
155
+ );
156
+
157
+ if (!meta) {
158
+ // Connected but not loaded (not persisted) - close it
159
+ //
160
+ // This may happen if the metadata was not successfully persisted
161
+ this.log?.warn({
162
+ msg: "closing websocket that is not persisted",
163
+ requestId: requestIdStr,
164
+ });
165
+
166
+ this.#sendMessage(gatewayId, requestId, {
167
+ tag: "ToServerWebSocketClose",
168
+ val: {
169
+ code: 1000,
170
+ reason: "ws.meta_not_found_during_restore",
171
+ hibernate: false,
172
+ },
173
+ });
174
+
175
+ connectedButNotLoadedCount++;
176
+ } else {
177
+ // Both connected and persisted - restore it
178
+ const request = buildRequestForWebSocket(
179
+ meta.path,
180
+ meta.headers,
181
+ );
182
+
183
+ // This will call `runner.config.websocket` under the hood to
184
+ // attach the event listeners to the WebSocket.
185
+ // Track this operation to ensure it completes
186
+ const restoreOperation = this.#createWebSocket(
187
+ actorId,
188
+ gatewayId,
189
+ requestId,
190
+ requestIdStr,
191
+ meta.serverMessageIndex,
192
+ true,
193
+ true,
194
+ request,
195
+ meta.path,
196
+ meta.headers,
197
+ false,
198
+ )
199
+ .then(() => {
200
+ // Create a PendingRequest entry to track the message index
201
+ const actor = this.#runner.getActor(actorId);
202
+ if (actor) {
203
+ actor.createPendingRequest(
204
+ gatewayId,
205
+ requestId,
206
+ meta.clientMessageIndex,
207
+ );
208
+ }
209
+
210
+ this.log?.info({
211
+ msg: "connection successfully restored",
212
+ actorId,
213
+ requestId: requestIdStr,
214
+ });
215
+ })
216
+ .catch((err) => {
217
+ this.log?.error({
218
+ msg: "error creating websocket during restore",
219
+ requestId: requestIdStr,
220
+ error: stringifyError(err),
221
+ });
222
+
223
+ // Close the WebSocket on error
224
+ this.#sendMessage(gatewayId, requestId, {
225
+ tag: "ToServerWebSocketClose",
226
+ val: {
227
+ code: 1011,
228
+ reason: "ws.restore_error",
229
+ hibernate: false,
230
+ },
231
+ });
232
+ });
233
+
234
+ backgroundOperations.push(restoreOperation);
235
+ restoredCount++;
236
+ }
237
+ }
238
+
239
+ // Process loaded but not connected (stale) - remove them
240
+ let loadedButNotConnectedCount = 0;
241
+ for (const meta of metaEntries) {
242
+ const requestIdStr = idToStr(meta.requestId);
243
+ const isConnected = actor.hibernatingRequests.some(
244
+ (req) =>
245
+ arraysEqual(req.gatewayId, meta.gatewayId) &&
246
+ arraysEqual(req.requestId, meta.requestId),
247
+ );
248
+ if (!isConnected) {
249
+ this.log?.warn({
250
+ msg: "removing stale persisted websocket",
251
+ requestId: requestIdStr,
252
+ });
253
+
254
+ const request = buildRequestForWebSocket(
255
+ meta.path,
256
+ meta.headers,
257
+ );
258
+
259
+ // Create adapter to register user's event listeners.
260
+ // Pass engineAlreadyClosed=true so close callback won't send tunnel message.
261
+ // Track this operation to ensure it completes
262
+ const cleanupOperation = this.#createWebSocket(
263
+ actorId,
264
+ meta.gatewayId,
265
+ meta.requestId,
266
+ requestIdStr,
267
+ meta.serverMessageIndex,
268
+ true,
269
+ true,
270
+ request,
271
+ meta.path,
272
+ meta.headers,
273
+ true,
274
+ )
275
+ .then((adapter) => {
276
+ // Close the adapter normally - this will fire user's close event handler
277
+ // (which should clean up persistence) and trigger the close callback
278
+ // (which will clean up maps but skip sending tunnel message)
279
+ adapter.close(1000, "ws.stale_metadata");
280
+ })
281
+ .catch((err) => {
282
+ this.log?.error({
283
+ msg: "error creating stale websocket during restore",
284
+ requestId: requestIdStr,
285
+ error: stringifyError(err),
286
+ });
287
+ });
288
+
289
+ backgroundOperations.push(cleanupOperation);
290
+ loadedButNotConnectedCount++;
291
+ }
292
+ }
293
+
294
+ // Wait for all background operations to complete before finishing
295
+ await Promise.allSettled(backgroundOperations);
296
+
297
+ // Mark restoration as complete
298
+ actor.hibernationRestored = true;
299
+
300
+ this.log?.info({
301
+ msg: "restored hibernatable websockets",
302
+ actorId,
303
+ restoredCount,
304
+ connectedButNotLoadedCount,
305
+ loadedButNotConnectedCount,
306
+ });
307
+ }
308
+
309
+ /**
310
+ * Called from WebSocketOpen message and when restoring hibernatable WebSockets.
311
+ *
312
+ * engineAlreadyClosed will be true if this is only being called to trigger
313
+ * the close callback and not to send a close message to the server. This
314
+ * is used specifically to clean up zombie WebSocket connections.
315
+ */
316
+ async #createWebSocket(
317
+ actorId: string,
318
+ gatewayId: GatewayId,
319
+ requestId: RequestId,
320
+ requestIdStr: string,
321
+ serverMessageIndex: number,
322
+ isHibernatable: boolean,
323
+ isRestoringHibernatable: boolean,
324
+ request: Request,
325
+ path: string,
326
+ headers: Record<string, string>,
327
+ engineAlreadyClosed: boolean,
328
+ ): Promise<WebSocketTunnelAdapter> {
329
+ this.log?.debug({
330
+ msg: "createWebSocket creating adapter",
331
+ actorId,
332
+ requestIdStr,
333
+ isHibernatable,
334
+ path,
335
+ });
336
+ // Create WebSocket adapter
337
+ const adapter = new WebSocketTunnelAdapter(
338
+ this,
339
+ actorId,
340
+ requestIdStr,
341
+ serverMessageIndex,
342
+ isHibernatable,
343
+ isRestoringHibernatable,
344
+ request,
345
+ (data: ArrayBuffer | string, isBinary: boolean) => {
346
+ // Send message through tunnel
347
+ const dataBuffer =
348
+ typeof data === "string"
349
+ ? (new TextEncoder().encode(data).buffer as ArrayBuffer)
350
+ : data;
351
+
352
+ this.#sendMessage(gatewayId, requestId, {
353
+ tag: "ToServerWebSocketMessage",
354
+ val: {
355
+ data: dataBuffer,
356
+ binary: isBinary,
357
+ },
358
+ });
359
+ },
360
+ (code?: number, reason?: string) => {
361
+ // Send close through tunnel if engine doesn't already know it's closed
362
+ if (!engineAlreadyClosed) {
363
+ this.#sendMessage(gatewayId, requestId, {
364
+ tag: "ToServerWebSocketClose",
365
+ val: {
366
+ code: code || null,
367
+ reason: reason || null,
368
+ hibernate: false,
369
+ },
370
+ });
371
+ }
372
+
373
+ // Clean up actor tracking
374
+ const actor = this.#runner.getActor(actorId);
375
+ if (actor) {
376
+ actor.deleteWebSocket(gatewayId, requestId);
377
+ actor.deletePendingRequest(gatewayId, requestId);
378
+ }
379
+
380
+ // Clean up request-to-actor mapping
381
+ this.#removeRequestToActor(gatewayId, requestId);
382
+ },
383
+ );
384
+
385
+ // Get actor and add websocket to it
386
+ const actor = this.#runner.getActor(actorId);
387
+ if (!actor) {
388
+ throw new Error(`Actor ${actorId} not found`);
389
+ }
390
+
391
+ actor.setWebSocket(gatewayId, requestId, adapter);
392
+ this.addRequestToActor(gatewayId, requestId, actorId);
393
+
394
+ // Call WebSocket handler. This handler will add event listeners
395
+ // for `open`, etc. Pass the VirtualWebSocket (not the adapter) to the actor.
396
+ await this.#runner.config.websocket(
397
+ this.#runner,
398
+ actorId,
399
+ adapter.websocket,
400
+ gatewayId,
401
+ requestId,
402
+ request,
403
+ path,
404
+ headers,
405
+ isHibernatable,
406
+ isRestoringHibernatable,
407
+ );
408
+
409
+ return adapter;
410
+ }
411
+
412
+ addRequestToActor(
413
+ gatewayId: GatewayId,
414
+ requestId: RequestId,
415
+ actorId: string,
416
+ ) {
417
+ this.#requestToActor.push({ gatewayId, requestId, actorId });
418
+ }
419
+
420
+ #removeRequestToActor(gatewayId: GatewayId, requestId: RequestId) {
421
+ const index = this.#requestToActor.findIndex(
422
+ (entry) =>
423
+ arraysEqual(entry.gatewayId, gatewayId) &&
424
+ arraysEqual(entry.requestId, requestId),
425
+ );
426
+ if (index !== -1) {
427
+ this.#requestToActor.splice(index, 1);
428
+ }
429
+ }
430
+
431
+ getRequestActor(
432
+ gatewayId: GatewayId,
433
+ requestId: RequestId,
434
+ ): RunnerActor | undefined {
435
+ const entry = this.#requestToActor.find(
436
+ (entry) =>
437
+ arraysEqual(entry.gatewayId, gatewayId) &&
438
+ arraysEqual(entry.requestId, requestId),
439
+ );
440
+
441
+ if (!entry) {
442
+ this.log?.warn({
443
+ msg: "missing requestToActor entry",
444
+ requestId: idToStr(requestId),
445
+ });
446
+ return undefined;
447
+ }
448
+
449
+ const actor = this.#runner.getActor(entry.actorId);
450
+ if (!actor) {
451
+ this.log?.warn({
452
+ msg: "missing actor for requestToActor lookup",
453
+ requestId: idToStr(requestId),
454
+ actorId: entry.actorId,
455
+ });
456
+ return undefined;
457
+ }
458
+
459
+ return actor;
460
+ }
461
+
462
+ async getAndWaitForRequestActor(
463
+ gatewayId: GatewayId,
464
+ requestId: RequestId,
465
+ ): Promise<RunnerActor | undefined> {
466
+ const actor = this.getRequestActor(gatewayId, requestId);
467
+ if (!actor) return;
468
+ await actor.actorStartPromise.promise;
469
+ return actor;
470
+ }
471
+
472
+ #sendMessage(
473
+ gatewayId: GatewayId,
474
+ requestId: RequestId,
475
+ messageKind: protocol.ToServerTunnelMessageKind,
476
+ ) {
477
+ // Buffer message if not connected
478
+ if (!this.#runner.getPegboardWebSocketIfReady()) {
479
+ this.log?.debug({
480
+ msg: "buffering tunnel message, socket not connected to engine",
481
+ requestId: idToStr(requestId),
482
+ message: stringifyToServerTunnelMessageKind(messageKind),
483
+ });
484
+ this.#bufferedMessages.push({ gatewayId, requestId, messageKind });
485
+ return;
486
+ }
487
+
488
+ // Get or initialize message index for this request
489
+ //
490
+ // We don't have to wait for the actor to start since we're not calling
491
+ // any callbacks on the actor
492
+ const gatewayIdStr = idToStr(gatewayId);
493
+ const requestIdStr = idToStr(requestId);
494
+ const actor = this.getRequestActor(gatewayId, requestId);
495
+ if (!actor) {
496
+ this.log?.warn({
497
+ msg: "cannot send tunnel message, actor not found",
498
+ gatewayId: gatewayIdStr,
499
+ requestId: requestIdStr,
500
+ });
501
+ return;
502
+ }
503
+
504
+ // Get message index from pending request
505
+ let clientMessageIndex: number;
506
+ const pending = actor.getPendingRequest(gatewayId, requestId);
507
+ if (pending) {
508
+ clientMessageIndex = pending.clientMessageIndex;
509
+ pending.clientMessageIndex++;
510
+ } else {
511
+ // No pending request
512
+ this.log?.warn({
513
+ msg: "missing pending request for send message, defaulting to message index 0",
514
+ gatewayId: gatewayIdStr,
515
+ requestId: requestIdStr,
516
+ });
517
+ clientMessageIndex = 0;
518
+ }
519
+
520
+ // Build message ID from gatewayId + requestId + messageIndex
521
+ const messageId: protocol.MessageId = {
522
+ gatewayId,
523
+ requestId,
524
+ messageIndex: clientMessageIndex,
525
+ };
526
+ const messageIdStr = `${idToStr(messageId.gatewayId)}-${idToStr(messageId.requestId)}-${messageId.messageIndex}`;
527
+
528
+ this.log?.debug({
529
+ msg: "sending tunnel msg",
530
+ messageId: messageIdStr,
531
+ gatewayId: gatewayIdStr,
532
+ requestId: requestIdStr,
533
+ messageIndex: clientMessageIndex,
534
+ message: stringifyToServerTunnelMessageKind(messageKind),
535
+ });
536
+
537
+ // Send message
538
+ const message: protocol.ToServer = {
539
+ tag: "ToServerTunnelMessage",
540
+ val: {
541
+ messageId,
542
+ messageKind,
543
+ },
544
+ };
545
+ this.#runner.__sendToServer(message);
546
+ }
547
+
548
+ closeActiveRequests(actor: RunnerActor) {
549
+ const actorId = actor.actorId;
550
+
551
+ // Terminate all requests for this actor. This will no send a
552
+ // ToServerResponse* message since the actor will no longer be loaded.
553
+ // The gateway is responsible for closing the request.
554
+ for (const entry of actor.pendingRequests) {
555
+ entry.request.reject(new Error(`Actor ${actorId} stopped`));
556
+ if (entry.gatewayId && entry.requestId) {
557
+ this.#removeRequestToActor(entry.gatewayId, entry.requestId);
558
+ }
559
+ }
560
+
561
+ // Close all WebSockets. Only send close event to non-HWS. The gateway is
562
+ // responsible for hibernating HWS and closing regular WS.
563
+ for (const entry of actor.webSockets) {
564
+ const isHibernatable = entry.ws[HIBERNATABLE_SYMBOL];
565
+ if (!isHibernatable) {
566
+ entry.ws._closeWithoutCallback(1000, "actor.stopped");
567
+ }
568
+ // Note: request-to-actor mapping is cleaned up in the close callback
569
+ }
570
+ }
571
+
572
+ async #fetch(
573
+ actorId: string,
574
+ gatewayId: protocol.GatewayId,
575
+ requestId: protocol.RequestId,
576
+ request: Request,
577
+ ): Promise<Response> {
578
+ // Validate actor exists
579
+ if (!this.#runner.hasActor(actorId)) {
580
+ this.log?.warn({
581
+ msg: "ignoring request for unknown actor",
582
+ actorId,
583
+ });
584
+
585
+ // NOTE: This is a special response that will cause Guard to retry the request
586
+ //
587
+ // See should_retry_request_inner
588
+ // https://github.com/rivet-dev/rivet/blob/222dae87e3efccaffa2b503de40ecf8afd4e31eb/engine/packages/guard-core/src/proxy_service.rs#L2458
589
+ return new Response("Actor not found", {
590
+ status: 503,
591
+ headers: { "x-rivet-error": "runner.actor_not_found" },
592
+ });
593
+ }
594
+
595
+ const fetchHandler = this.#runner.config.fetch(
596
+ this.#runner,
597
+ actorId,
598
+ gatewayId,
599
+ requestId,
600
+ request,
601
+ );
602
+
603
+ if (!fetchHandler) {
604
+ return new Response("Not Implemented", { status: 501 });
605
+ }
606
+
607
+ return fetchHandler;
608
+ }
609
+
610
+ async handleTunnelMessage(message: protocol.ToClientTunnelMessage) {
611
+ // Parse the gateway ID, request ID, and message index from the messageId
612
+ const { gatewayId, requestId, messageIndex } = message.messageId;
613
+
614
+ const gatewayIdStr = idToStr(gatewayId);
615
+ const requestIdStr = idToStr(requestId);
616
+ this.log?.debug({
617
+ msg: "receive tunnel msg",
618
+ gatewayId: gatewayIdStr,
619
+ requestId: requestIdStr,
620
+ messageIndex: message.messageId.messageIndex,
621
+ message: stringifyToClientTunnelMessageKind(message.messageKind),
622
+ });
623
+
624
+ switch (message.messageKind.tag) {
625
+ case "ToClientRequestStart":
626
+ await this.#handleRequestStart(
627
+ gatewayId,
628
+ requestId,
629
+ message.messageKind.val,
630
+ );
631
+ break;
632
+ case "ToClientRequestChunk":
633
+ await this.#handleRequestChunk(
634
+ gatewayId,
635
+ requestId,
636
+ message.messageKind.val,
637
+ );
638
+ break;
639
+ case "ToClientRequestAbort":
640
+ await this.#handleRequestAbort(gatewayId, requestId);
641
+ break;
642
+ case "ToClientWebSocketOpen":
643
+ await this.#handleWebSocketOpen(
644
+ gatewayId,
645
+ requestId,
646
+ message.messageKind.val,
647
+ );
648
+ break;
649
+ case "ToClientWebSocketMessage": {
650
+ await this.#handleWebSocketMessage(
651
+ gatewayId,
652
+ requestId,
653
+ messageIndex,
654
+ message.messageKind.val,
655
+ );
656
+ break;
657
+ }
658
+ case "ToClientWebSocketClose":
659
+ await this.#handleWebSocketClose(
660
+ gatewayId,
661
+ requestId,
662
+ message.messageKind.val,
663
+ );
664
+ break;
665
+ default:
666
+ unreachable(message.messageKind);
667
+ }
668
+ }
669
+
670
+ async #handleRequestStart(
671
+ gatewayId: GatewayId,
672
+ requestId: RequestId,
673
+ req: protocol.ToClientRequestStart,
674
+ ) {
675
+ // Track this request for the actor
676
+ const requestIdStr = idToStr(requestId);
677
+ const actor = await this.#runner.getAndWaitForActor(req.actorId);
678
+ if (!actor) {
679
+ this.log?.warn({
680
+ msg: "actor does not exist in handleRequestStart, request will leak",
681
+ actorId: req.actorId,
682
+ requestId: requestIdStr,
683
+ });
684
+ return;
685
+ }
686
+
687
+ // Add to request-to-actor mapping
688
+ this.addRequestToActor(gatewayId, requestId, req.actorId);
689
+
690
+ try {
691
+ // Convert headers map to Headers object
692
+ const headers = new Headers();
693
+ for (const [key, value] of req.headers) {
694
+ headers.append(key, value);
695
+ }
696
+
697
+ // Create Request object
698
+ const request = new Request(`http://localhost${req.path}`, {
699
+ method: req.method,
700
+ headers,
701
+ body: req.body ? new Uint8Array(req.body) : undefined,
702
+ });
703
+
704
+ // Handle streaming request
705
+ if (req.stream) {
706
+ // Create a stream for the request body
707
+ const stream = new ReadableStream<Uint8Array>({
708
+ start: (controller) => {
709
+ // Store controller for chunks
710
+ const existing = actor.getPendingRequest(
711
+ gatewayId,
712
+ requestId,
713
+ );
714
+ if (existing) {
715
+ existing.streamController = controller;
716
+ existing.actorId = req.actorId;
717
+ existing.gatewayId = gatewayId;
718
+ existing.requestId = requestId;
719
+ } else {
720
+ actor.createPendingRequestWithStreamController(
721
+ gatewayId,
722
+ requestId,
723
+ 0,
724
+ controller,
725
+ );
726
+ }
727
+ },
728
+ });
729
+
730
+ // Create request with streaming body
731
+ const streamingRequest = new Request(request, {
732
+ body: stream,
733
+ duplex: "half",
734
+ } as any);
735
+
736
+ // Call fetch handler with validation
737
+ const response = await this.#fetch(
738
+ req.actorId,
739
+ gatewayId,
740
+ requestId,
741
+ streamingRequest,
742
+ );
743
+ await this.#sendResponse(
744
+ actor.actorId,
745
+ actor.generation,
746
+ gatewayId,
747
+ requestId,
748
+ response,
749
+ );
750
+ } else {
751
+ // Non-streaming request
752
+ // Create a pending request entry to track messageIndex for the response
753
+ actor.createPendingRequest(gatewayId, requestId, 0);
754
+
755
+ const response = await this.#fetch(
756
+ req.actorId,
757
+ gatewayId,
758
+ requestId,
759
+ request,
760
+ );
761
+ await this.#sendResponse(
762
+ actor.actorId,
763
+ actor.generation,
764
+ gatewayId,
765
+ requestId,
766
+ response,
767
+ );
768
+ }
769
+ } catch (error) {
770
+ if (error instanceof RunnerShutdownError) {
771
+ this.log?.debug({ msg: "catught runner shutdown error" });
772
+ } else {
773
+ this.log?.error({ msg: "error handling request", error });
774
+ this.#sendResponseError(
775
+ actor.actorId,
776
+ actor.generation,
777
+ gatewayId,
778
+ requestId,
779
+ 500,
780
+ "Internal Server Error",
781
+ );
782
+ }
783
+ } finally {
784
+ // Clean up request tracking
785
+ if (this.#runner.hasActor(req.actorId, actor.generation)) {
786
+ actor.deletePendingRequest(gatewayId, requestId);
787
+ this.#removeRequestToActor(gatewayId, requestId);
788
+ }
789
+ }
790
+ }
791
+
792
+ async #handleRequestChunk(
793
+ gatewayId: GatewayId,
794
+ requestId: RequestId,
795
+ chunk: protocol.ToClientRequestChunk,
796
+ ) {
797
+ const actor = await this.getAndWaitForRequestActor(
798
+ gatewayId,
799
+ requestId,
800
+ );
801
+ if (actor) {
802
+ const pending = actor.getPendingRequest(gatewayId, requestId);
803
+ if (pending?.streamController) {
804
+ pending.streamController.enqueue(new Uint8Array(chunk.body));
805
+ if (chunk.finish) {
806
+ pending.streamController.close();
807
+ actor.deletePendingRequest(gatewayId, requestId);
808
+ this.#removeRequestToActor(gatewayId, requestId);
809
+ }
810
+ }
811
+ }
812
+ }
813
+
814
+ async #handleRequestAbort(gatewayId: GatewayId, requestId: RequestId) {
815
+ const actor = await this.getAndWaitForRequestActor(
816
+ gatewayId,
817
+ requestId,
818
+ );
819
+ if (actor) {
820
+ const pending = actor.getPendingRequest(gatewayId, requestId);
821
+ if (pending?.streamController) {
822
+ pending.streamController.error(new Error("Request aborted"));
823
+ }
824
+ actor.deletePendingRequest(gatewayId, requestId);
825
+ this.#removeRequestToActor(gatewayId, requestId);
826
+ }
827
+ }
828
+
829
+ async #sendResponse(
830
+ actorId: string,
831
+ generation: number,
832
+ gatewayId: GatewayId,
833
+ requestId: ArrayBuffer,
834
+ response: Response,
835
+ ) {
836
+ if (!this.#runner.hasActor(actorId, generation)) {
837
+ this.log?.warn({
838
+ msg: "actor not loaded to send response, assuming gateway has closed request",
839
+ actorId,
840
+ generation,
841
+ requestId,
842
+ });
843
+ return;
844
+ }
845
+
846
+ // Always treat responses as non-streaming for now
847
+ // In the future, we could detect streaming responses based on:
848
+ // - Transfer-Encoding: chunked
849
+ // - Content-Type: text/event-stream
850
+ // - Explicit stream flag from the handler
851
+
852
+ // Read the body first to get the actual content
853
+ const body = response.body ? await response.arrayBuffer() : null;
854
+
855
+ if (body && body.byteLength > MAX_PAYLOAD_SIZE) {
856
+ throw new Error("Response body too large");
857
+ }
858
+
859
+ // Convert headers to map and add Content-Length if not present
860
+ const headers = new Map<string, string>();
861
+ response.headers.forEach((value, key) => {
862
+ headers.set(key, value);
863
+ });
864
+
865
+ // Add Content-Length header if we have a body and it's not already set
866
+ if (body && !headers.has("content-length")) {
867
+ headers.set("content-length", String(body.byteLength));
868
+ }
869
+
870
+ // Send as non-streaming response if actor has not stopped
871
+ this.#sendMessage(gatewayId, requestId, {
872
+ tag: "ToServerResponseStart",
873
+ val: {
874
+ status: response.status as protocol.u16,
875
+ headers,
876
+ body: body || null,
877
+ stream: false,
878
+ },
879
+ });
880
+ }
881
+
882
+ #sendResponseError(
883
+ actorId: string,
884
+ generation: number,
885
+ gatewayId: GatewayId,
886
+ requestId: ArrayBuffer,
887
+ status: number,
888
+ message: string,
889
+ ) {
890
+ if (!this.#runner.hasActor(actorId, generation)) {
891
+ this.log?.warn({
892
+ msg: "actor not loaded to send response, assuming gateway has closed request",
893
+ actorId,
894
+ generation,
895
+ requestId,
896
+ });
897
+ return;
898
+ }
899
+
900
+ const headers = new Map<string, string>();
901
+ headers.set("content-type", "text/plain");
902
+
903
+ this.#sendMessage(gatewayId, requestId, {
904
+ tag: "ToServerResponseStart",
905
+ val: {
906
+ status: status as protocol.u16,
907
+ headers,
908
+ body: new TextEncoder().encode(message).buffer as ArrayBuffer,
909
+ stream: false,
910
+ },
911
+ });
912
+ }
913
+
914
+ async #handleWebSocketOpen(
915
+ gatewayId: GatewayId,
916
+ requestId: RequestId,
917
+ open: protocol.ToClientWebSocketOpen,
918
+ ) {
919
+ // NOTE: This method is safe to be async since we will not receive any
920
+ // further WebSocket events until we send a ToServerWebSocketOpen
921
+ // tunnel message. We can do any async logic we need to between those two events.
922
+ //
923
+ // Sending a ToServerWebSocketClose will terminate the WebSocket early.
924
+
925
+ const requestIdStr = idToStr(requestId);
926
+
927
+ // Validate actor exists
928
+ const actor = await this.#runner.getAndWaitForActor(open.actorId);
929
+ if (!actor) {
930
+ this.log?.warn({
931
+ msg: "ignoring websocket for unknown actor",
932
+ actorId: open.actorId,
933
+ });
934
+
935
+ // NOTE: Closing a WebSocket before open is equivalent to a Service
936
+ // Unavailable error and will cause Guard to retry the request
937
+ //
938
+ // See
939
+ // https://github.com/rivet-dev/rivet/blob/222dae87e3efccaffa2b503de40ecf8afd4e31eb/engine/packages/pegboard-gateway/src/lib.rs#L238
940
+ this.#sendMessage(gatewayId, requestId, {
941
+ tag: "ToServerWebSocketClose",
942
+ val: {
943
+ code: 1011,
944
+ reason: "Actor not found",
945
+ hibernate: false,
946
+ },
947
+ });
948
+ return;
949
+ }
950
+
951
+ // Close existing WebSocket if one already exists for this request ID.
952
+ // This should never happen, but prevents any potential duplicate
953
+ // WebSockets from retransmits.
954
+ const existingAdapter = actor.getWebSocket(gatewayId, requestId);
955
+ if (existingAdapter) {
956
+ this.log?.warn({
957
+ msg: "closing existing websocket for duplicate open event for the same request id",
958
+ requestId: requestIdStr,
959
+ });
960
+ // Close without sending a message through the tunnel since the server
961
+ // already knows about the new connection
962
+ existingAdapter._closeWithoutCallback(1000, "ws.duplicate_open");
963
+ }
964
+
965
+ // Create WebSocket
966
+ try {
967
+ const request = buildRequestForWebSocket(
968
+ open.path,
969
+ Object.fromEntries(open.headers),
970
+ );
971
+
972
+ const canHibernate =
973
+ this.#runner.config.hibernatableWebSocket.canHibernate(
974
+ actor.actorId,
975
+ gatewayId,
976
+ requestId,
977
+ request,
978
+ );
979
+
980
+ // #createWebSocket will call `runner.config.websocket` under the
981
+ // hood to add the event listeners for open, etc. If this handler
982
+ // throws, then the WebSocket will be closed before sending the
983
+ // open event.
984
+ const adapter = await this.#createWebSocket(
985
+ actor.actorId,
986
+ gatewayId,
987
+ requestId,
988
+ requestIdStr,
989
+ 0,
990
+ canHibernate,
991
+ false,
992
+ request,
993
+ open.path,
994
+ Object.fromEntries(open.headers),
995
+ false,
996
+ );
997
+
998
+ // Create a PendingRequest entry to track the message index
999
+ actor.createPendingRequest(gatewayId, requestId, 0);
1000
+
1001
+ // Open the WebSocket after `config.socket` so (a) the event
1002
+ // handlers can be added and (b) any errors in `config.websocket`
1003
+ // will cause the WebSocket to terminate before the open event.
1004
+ this.#sendMessage(gatewayId, requestId, {
1005
+ tag: "ToServerWebSocketOpen",
1006
+ val: {
1007
+ canHibernate,
1008
+ },
1009
+ });
1010
+
1011
+ // Dispatch open event
1012
+ adapter._handleOpen(requestId);
1013
+ } catch (error) {
1014
+ this.log?.error({ msg: "error handling websocket open", error });
1015
+
1016
+ // TODO: Call close event on adapter if needed
1017
+
1018
+ // Send close on error
1019
+ this.#sendMessage(gatewayId, requestId, {
1020
+ tag: "ToServerWebSocketClose",
1021
+ val: {
1022
+ code: 1011,
1023
+ reason: "Server Error",
1024
+ hibernate: false,
1025
+ },
1026
+ });
1027
+
1028
+ // Clean up actor tracking
1029
+ actor.deleteWebSocket(gatewayId, requestId);
1030
+ actor.deletePendingRequest(gatewayId, requestId);
1031
+ this.#removeRequestToActor(gatewayId, requestId);
1032
+ }
1033
+ }
1034
+
1035
+ async #handleWebSocketMessage(
1036
+ gatewayId: GatewayId,
1037
+ requestId: RequestId,
1038
+ serverMessageIndex: number,
1039
+ msg: protocol.ToClientWebSocketMessage,
1040
+ ) {
1041
+ const actor = await this.getAndWaitForRequestActor(
1042
+ gatewayId,
1043
+ requestId,
1044
+ );
1045
+ if (actor) {
1046
+ const adapter = actor.getWebSocket(gatewayId, requestId);
1047
+ if (adapter) {
1048
+ const data = msg.binary
1049
+ ? new Uint8Array(msg.data)
1050
+ : new TextDecoder().decode(new Uint8Array(msg.data));
1051
+
1052
+ adapter._handleMessage(
1053
+ requestId,
1054
+ data,
1055
+ serverMessageIndex,
1056
+ msg.binary,
1057
+ );
1058
+ return;
1059
+ }
1060
+ }
1061
+
1062
+ // TODO: This will never retransmit the socket and the socket will close
1063
+ this.log?.warn({
1064
+ msg: "missing websocket for incoming websocket message, this may indicate the actor stopped before processing a message",
1065
+ requestId,
1066
+ });
1067
+ }
1068
+
1069
+ sendHibernatableWebSocketMessageAck(
1070
+ gatewayId: ArrayBuffer,
1071
+ requestId: ArrayBuffer,
1072
+ clientMessageIndex: number,
1073
+ ) {
1074
+ const requestIdStr = idToStr(requestId);
1075
+
1076
+ this.log?.debug({
1077
+ msg: "ack ws msg",
1078
+ requestId: requestIdStr,
1079
+ index: clientMessageIndex,
1080
+ });
1081
+
1082
+ if (clientMessageIndex < 0 || clientMessageIndex > 65535)
1083
+ throw new Error("Invalid websocket ack index");
1084
+
1085
+ // Get the actor to find the gatewayId
1086
+ //
1087
+ // We don't have to wait for the actor to start since we're not calling
1088
+ // any callbacks on the actor
1089
+ const actor = this.getRequestActor(gatewayId, requestId);
1090
+ if (!actor) {
1091
+ this.log?.warn({
1092
+ msg: "cannot send websocket ack, actor not found",
1093
+ requestId: requestIdStr,
1094
+ });
1095
+ return;
1096
+ }
1097
+
1098
+ // Get gatewayId from the pending request
1099
+ const pending = actor.getPendingRequest(gatewayId, requestId);
1100
+ if (!pending?.gatewayId) {
1101
+ this.log?.warn({
1102
+ msg: "cannot send websocket ack, gatewayId not found in pending request",
1103
+ requestId: requestIdStr,
1104
+ });
1105
+ return;
1106
+ }
1107
+
1108
+ // Send the ack message
1109
+ this.#sendMessage(pending.gatewayId, requestId, {
1110
+ tag: "ToServerWebSocketMessageAck",
1111
+ val: {
1112
+ index: clientMessageIndex,
1113
+ },
1114
+ });
1115
+ }
1116
+
1117
+ async #handleWebSocketClose(
1118
+ gatewayId: GatewayId,
1119
+ requestId: RequestId,
1120
+ close: protocol.ToClientWebSocketClose,
1121
+ ) {
1122
+ const actor = await this.getAndWaitForRequestActor(
1123
+ gatewayId,
1124
+ requestId,
1125
+ );
1126
+ if (actor) {
1127
+ const adapter = actor.getWebSocket(gatewayId, requestId);
1128
+ if (adapter) {
1129
+ // We don't need to send a close response
1130
+ adapter._handleClose(
1131
+ requestId,
1132
+ close.code || undefined,
1133
+ close.reason || undefined,
1134
+ );
1135
+ actor.deleteWebSocket(gatewayId, requestId);
1136
+ actor.deletePendingRequest(gatewayId, requestId);
1137
+ this.#removeRequestToActor(gatewayId, requestId);
1138
+ }
1139
+ }
1140
+ }
1141
+ }
1142
+
1143
+ /**
1144
+ * Builds a request that represents the incoming request for a given WebSocket.
1145
+ *
1146
+ * This request is not a real request and will never be sent. It's used to be passed to the actor to behave like a real incoming request.
1147
+ */
1148
+ function buildRequestForWebSocket(
1149
+ path: string,
1150
+ headers: Record<string, string>,
1151
+ ): Request {
1152
+ // We need to manually ensure the original Upgrade/Connection WS
1153
+ // headers are present
1154
+ const fullHeaders = {
1155
+ ...headers,
1156
+ Upgrade: "websocket",
1157
+ Connection: "Upgrade",
1158
+ };
1159
+
1160
+ if (!path.startsWith("/")) {
1161
+ throw new Error("Path must start with leading slash");
1162
+ }
1163
+
1164
+ const request = new Request(`http://actor${path}`, {
1165
+ method: "GET",
1166
+ headers: fullHeaders,
1167
+ });
1168
+
1169
+ return request;
1170
+ }