import z from "zod"; import type { RpcContract, RpcContractDefinition, RpcProcedureDefinition } from "./contract.ts"; import type { RpcTransport } from "./transports/rpc-transport.ts"; import type { Message } from "./messages.ts"; import { PendingResponses } from "./pending-responses.ts"; import { v7 as uuidv7 } from "uuid"; type TimeoutId = ReturnType; export type RpcPeerOptions = { transport: RpcTransport; localContract: RpcContract; remoteContract: RpcContract; }; type MaybeAsyncGenerator = | AsyncGenerator | Generator; type MaybePromise = T | Promise; type ProcedureResult = T_Definition["generator"] extends true ? MaybeAsyncGenerator, void, unknown> : MaybePromise>; type LocalHandlers = { [Method in keyof T_Definition]: ( params: z.infer, ) => ProcedureResult; }; type RemoteCallReturn = T_Definition["generator"] extends true ? AsyncGenerator, void, unknown> : Awaited>; const STREAM_IDLE_TIMEOUT_MS = 30000; type StreamEntry = { streamId: string; generator: MaybeAsyncGenerator; timeoutId?: TimeoutId; }; export class RpcPeer { private _options: T_Options; private _handlers: LocalHandlers; private _pendingResponses: PendingResponses; private _streams: Map; constructor( options: T_Options, handlers: LocalHandlers, ) { this._options = options; this._handlers = handlers; this._pendingResponses = new PendingResponses(); this._streams = new Map(); this._options.transport.onMessage((message) => { this._handleIncomingMessage(message); }); } async call( method: T_Method, params: z.infer, ): Promise> { const msgId = uuidv7(); const response = await this._sendMessageAndWaitForResponse({ type: "callMethod", method: method as string, value: params, msgId, }); if (response.type === "callResponse") { return response.value as RemoteCallReturn< T_Options["remoteContract"]["definition"][T_Method] >; } else if (response.type === "callResponseStream") { return this._createStreamGenerator(response.streamId, response.senderId) as any; } else { throw new Error(`Unexpected message type: ${response.type}`); } } private async *_createStreamGenerator(streamId: string, senderId: string | undefined) { while (true) { const msgId = uuidv7(); const chunkResponse = await this._sendMessageAndWaitForResponse({ type: "requestStreamChunk", streamId, msgId, value: undefined, senderId, }); if (chunkResponse.type === "streamChunkValue") { yield chunkResponse.value; } else if (chunkResponse.type === "streamChunkDone") { return; } else if (chunkResponse.type === "streamChunkError") { throw new Error(`Stream error from remote: ${chunkResponse.error}`); } else { throw new Error(`Unexpected message type: ${chunkResponse.type}`); } } } private _refreshStreamTimeout(streamId: string) { const stream = this._streams.get(streamId); if (!stream) return; if (stream.timeoutId) { clearTimeout(stream.timeoutId); } stream.timeoutId = setTimeout(() => { this._streams.delete(streamId); }, STREAM_IDLE_TIMEOUT_MS); } private _deleteStream(streamId: string) { const stream = this._streams.get(streamId); if (!stream) return; if (stream.timeoutId) { clearTimeout(stream.timeoutId); } this._streams.delete(streamId); } private async _handleIncomingMessage(message: Message) { if ("prevMsgId" in message) { this._pendingResponses.resolve(message.prevMsgId, message); } if (message.type === "ping") { this._options.transport.sendMessage({ type: "pong", msgId: uuidv7(), prevMsgId: message.msgId, }); return; } else if (message.type === "pong") { return; } else if (message.type === "callMethod") { const definition = this._options.localContract.definition[ message.method as keyof T_Options["localContract"]["definition"] ]; if (!definition) { this._sendMessage({ type: "callResponse", value: undefined, msgId: uuidv7(), prevMsgId: message.msgId, senderId: message.senderId, error: `No method named ${message.method}`, }); return; } const handler = this._handlers[message.method as keyof T_Options["localContract"]["definition"]]; if (!handler) { this._sendMessage({ type: "callResponse", value: undefined, msgId: uuidv7(), prevMsgId: message.msgId, senderId: message.senderId, error: `No handler for method ${message.method}`, }); return; } try { const result = await handler(message.value); if (isGenerator(result) || isAsyncGenerator(result)) { const streamId = uuidv7(); this._streams.set(streamId, { streamId, generator: result }); this._refreshStreamTimeout(streamId); this._sendMessage({ type: "callResponseStream", msgId: uuidv7(), prevMsgId: message.msgId, senderId: message.senderId, streamId, }); } else { this._sendMessage({ type: "callResponse", msgId: uuidv7(), prevMsgId: message.msgId, senderId: message.senderId, value: result, }); } } catch (err) { this._sendMessage({ type: "callResponse", msgId: uuidv7(), value: undefined, prevMsgId: message.msgId, senderId: message.senderId, error: err instanceof Error ? err.message : String(err), }); } return; } else if (message.type === "requestStreamChunk") { const stream = this._streams.get(message.streamId); if (!stream) { this._sendMessage({ type: "streamChunkError", msgId: uuidv7(), prevMsgId: message.msgId, streamId: message.streamId, error: `No stream with id ${message.streamId}`, senderId: message.senderId, }); return; } try { const { value, done } = await stream.generator.next(); if (done) { this._sendMessage({ type: "streamChunkDone", msgId: uuidv7(), prevMsgId: message.msgId, streamId: message.streamId, senderId: message.senderId, }); this._deleteStream(message.streamId); } else { this._refreshStreamTimeout(message.streamId); this._sendMessage({ type: "streamChunkValue", msgId: uuidv7(), prevMsgId: message.msgId, streamId: message.streamId, senderId: message.senderId, value, }); } } catch (err) { this._options.transport.sendMessage({ type: "streamChunkError", msgId: uuidv7(), prevMsgId: message.msgId, streamId: message.streamId, error: err instanceof Error ? err.message : String(err), senderId: message.senderId, }); this._streams.delete(message.streamId); } return; } else if (message.type === "callResponse") { return; } } private _sendMessage(message: Message) { this._options.transport.sendMessage(message); } private async _sendMessageAndWaitForResponse(message: Message): Promise { this._options.transport.sendMessage(message); return this._pendingResponses.waitFor(message.msgId); } } function isGenerator(obj: any): obj is Generator { return ( obj && typeof obj.next === "function" && typeof obj.throw === "function" && typeof obj.return === "function" ); } function isAsyncGenerator(obj: any): obj is AsyncGenerator { return ( obj && typeof obj.next === "function" && typeof obj.throw === "function" && typeof obj.return === "function" && obj[Symbol.toStringTag] === "AsyncGenerator" ); }