diff --git a/packages/opencode/src/util/rpc.ts b/packages/opencode/src/util/rpc.ts index 02586ebcfc6..19d845a96df 100644 --- a/packages/opencode/src/util/rpc.ts +++ b/packages/opencode/src/util/rpc.ts @@ -6,8 +6,12 @@ export function listen(rpc: Definition) { onmessage = async (evt) => { const parsed = JSON.parse(evt.data) if (parsed.type === "rpc.request") { - const result = await rpc[parsed.method](parsed.input) - postMessage(JSON.stringify({ type: "rpc.result", result, id: parsed.id })) + try { + const result = await rpc[parsed.method](parsed.input) + postMessage(JSON.stringify({ type: "rpc.result", result, id: parsed.id })) + } catch (error) { + postMessage(JSON.stringify({ type: "rpc.error", error: serializeError(error), id: parsed.id })) + } } } } @@ -19,16 +23,35 @@ export function emit(event: string, data: unknown) { export function client(target: { postMessage: (data: string) => void | null onmessage: ((this: Worker, ev: MessageEvent) => any) | null + addEventListener?: Worker["addEventListener"] }) { - const pending = new Map void>() + const pending = new Map void; reject: (error: any) => void }>() const listeners = new Map void>>() + let failed: unknown let id = 0 + const rejectPending = (error: unknown) => { + failed = error + for (const request of pending.values()) { + request.reject(error) + } + pending.clear() + } + target.addEventListener?.("error", (event) => { + rejectPending(errorFromEvent(event)) + }) target.onmessage = async (evt) => { const parsed = JSON.parse(evt.data) if (parsed.type === "rpc.result") { - const resolve = pending.get(parsed.id) - if (resolve) { - resolve(parsed.result) + const request = pending.get(parsed.id) + if (request) { + request.resolve(parsed.result) + pending.delete(parsed.id) + } + } + if (parsed.type === "rpc.error") { + const request = pending.get(parsed.id) + if (request) { + request.reject(deserializeError(parsed.error)) pending.delete(parsed.id) } } @@ -43,10 +66,16 @@ export function client(target: { } return { call(method: Method, input: Parameters[0]): Promise> { + if (failed) return Promise.reject(failed) const requestId = id++ - return new Promise((resolve) => { - pending.set(requestId, resolve) - target.postMessage(JSON.stringify({ type: "rpc.request", method, input, id: requestId })) + return new Promise((resolve, reject) => { + pending.set(requestId, { resolve, reject }) + try { + target.postMessage(JSON.stringify({ type: "rpc.request", method, input, id: requestId })) + } catch (error) { + pending.delete(requestId) + reject(error) + } }) }, on(event: string, handler: (data: Data) => void) { @@ -63,4 +92,34 @@ export function client(target: { } } +function errorFromEvent(event: Event): unknown { + const errorEvent = event as { error?: unknown; message?: unknown } + if (errorEvent.error) return errorEvent.error + if (typeof errorEvent.message === "string" && errorEvent.message) return new Error(errorEvent.message) + return new Error("Worker failed") +} + +function serializeError(error: unknown): unknown { + if (!(error instanceof Error)) return error + return { + ...Object.fromEntries(Object.getOwnPropertyNames(error).map((key) => [key, error[key as keyof Error]])), + name: error.name, + message: error.message, + stack: error.stack, + cause: serializeError(error.cause), + } +} + +function deserializeError(input: unknown): unknown { + if (!input || typeof input !== "object" || !("message" in input)) return input + const serialized = input as { name?: unknown; message?: unknown; stack?: unknown; cause?: unknown } + const error = new Error(typeof serialized.message === "string" ? serialized.message : String(serialized.message), { + cause: deserializeError(serialized.cause), + }) + if (typeof serialized.name === "string") error.name = serialized.name + if (typeof serialized.stack === "string") error.stack = serialized.stack + Object.assign(error, input) + return error +} + export * as Rpc from "./rpc" diff --git a/packages/opencode/test/util/rpc.test.ts b/packages/opencode/test/util/rpc.test.ts new file mode 100644 index 00000000000..ba25531b09d --- /dev/null +++ b/packages/opencode/test/util/rpc.test.ts @@ -0,0 +1,50 @@ +import { describe, expect, test } from "bun:test" +import { Rpc } from "@/util/rpc" + +type TestRpc = { + fail(input: undefined): Promise +} + +type Target = Parameters>[0] + +describe("Rpc", () => { + test("rejects pending calls when the worker reports an error", async () => { + const target: Target = { + postMessage(data) { + const request = JSON.parse(data) + target.onmessage?.call( + {} as Worker, + { + data: JSON.stringify({ + type: "rpc.error", + id: request.id, + error: { name: "Error", message: "boom", stack: "Error: boom" }, + }), + } as MessageEvent, + ) + }, + onmessage: null, + } + + await expect(Rpc.client(target).call("fail", undefined)).rejects.toThrow("boom") + }) + + test("rejects pending and future calls when the worker crashes", async () => { + let onError: ((event: Event) => void) | undefined + const target: Target = { + postMessage() {}, + onmessage: null, + addEventListener(type, listener) { + if (type !== "error" || typeof listener !== "function") return + onError = (event) => listener.call({} as Worker, event) + }, + } + const client = Rpc.client(target) + const pending = client.call("fail", undefined) + + onError?.({ message: "worker crashed" } as Event) + + await expect(pending).rejects.toThrow("worker crashed") + await expect(client.call("fail", undefined)).rejects.toThrow("worker crashed") + }) +})