UNPKG

trpc-bun-adapter

Version:

TRPC adapter for bun js runtime

389 lines (387 loc) 10.3 kB
// src/createBunHttpHandler.ts import { fetchRequestHandler } from "@trpc/server/adapters/fetch"; function createBunHttpHandler(opts) { return (request, server) => { const url = new URL(request.url); if (opts.endpoint && !url.pathname.startsWith(opts.endpoint)) { return; } if (opts.emitWsUpgrades && server.upgrade(request, { data: { req: request } })) { return new Response(null, { status: 101 }); } return fetchRequestHandler({ createContext: () => ({}), ...opts, req: request, endpoint: opts.endpoint ?? "" }); }; } // src/createBunWSHandler.ts import { TRPCError, callTRPCProcedure, getErrorShape, getTRPCErrorFromUnknown, isTrackedEnvelope, transformTRPCResponse } from "@trpc/server"; import { parseConnectionParamsFromUnknown } from "@trpc/server/http"; import { isObservable, observableToAsyncIterable } from "@trpc/server/observable"; import { parseTRPCMessage } from "@trpc/server/rpc"; function createBunWSHandler(opts) { const { router, createContext } = opts; const respond = (client, untransformedJSON) => { client.send( JSON.stringify( transformTRPCResponse( opts.router._def._config, untransformedJSON ) ) ); }; async function createClientCtx(client, url, connectionParams) { const ctxPromise = createContext?.({ req: client.data.req, res: client, info: { url, connectionParams, calls: [], isBatchCall: false, accept: null, type: "unknown", signal: client.data.abortController.signal } }); try { return await ctxPromise; } catch (cause) { const error = getTRPCErrorFromUnknown(cause); opts.onError?.({ error, path: void 0, type: "unknown", ctx: void 0, req: client.data.req, input: void 0 }); respond(client, { id: null, error: getErrorShape({ config: router._def._config, error, type: "unknown", path: void 0, input: void 0, ctx: void 0 }) }); } } async function handleRequest(client, msg) { if (!msg.id) { throw new TRPCError({ code: "BAD_REQUEST", message: "`id` is required" }); } if (msg.method === "subscription.stop") { client.data.abortControllers.get(msg.id)?.abort(); client.data.abortControllers.delete(msg.id); return; } const { id, method, jsonrpc } = msg; const type = method; const { path, lastEventId } = msg.params; const req = client.data.req; const clientAbortControllers = client.data.abortControllers; let { input } = msg.params; const ctx = await client.data.ctx; try { if (lastEventId !== void 0) { if (isObject(input)) { input = { ...input, lastEventId }; } else { input ??= { lastEventId }; } } if (clientAbortControllers.has(id)) { throw new TRPCError({ message: `Duplicate id ${id}`, code: "BAD_REQUEST" }); } const abortController = new AbortController(); const result = await callTRPCProcedure({ router, path, getRawInput: () => Promise.resolve(input), ctx, type, signal: abortController.signal }); const isIterableResult = isAsyncIterable(result) || isObservable(result); if (type !== "subscription") { if (isIterableResult) { throw new TRPCError({ code: "UNSUPPORTED_MEDIA_TYPE", message: `Cannot return an async iterable or observable from a ${type} procedure with WebSockets` }); } respond(client, { id, jsonrpc, result: { type: "data", data: result } }); return; } if (!isIterableResult) { throw new TRPCError({ message: `Subscription ${path} did not return an observable or a AsyncGenerator`, code: "INTERNAL_SERVER_ERROR" }); } if (client.readyState !== WebSocket.OPEN) { return; } const iterable = isObservable(result) ? observableToAsyncIterable(result, abortController.signal) : result; const iterator = iterable[Symbol.asyncIterator](); const abortPromise = new Promise((resolve) => { abortController.signal.onabort = () => resolve("abort"); }); clientAbortControllers.set(id, abortController); run(async () => { while (true) { const next = await Promise.race([ iterator.next().catch(getTRPCErrorFromUnknown), abortPromise ]); if (next === "abort") { await iterator.return?.(); break; } if (next instanceof Error) { const error = getTRPCErrorFromUnknown(next); opts.onError?.({ error, path, type, ctx, req, input }); respond(client, { id, jsonrpc, error: getErrorShape({ config: router._def._config, error, type, path, input, ctx }) }); break; } if (next.done) { break; } const result2 = { type: "data", data: next.value }; if (isTrackedEnvelope(next.value)) { const [id2, data] = next.value; result2.id = id2; result2.data = { id: id2, data }; } respond(client, { id, jsonrpc, result: result2 }); } await iterator.return?.(); respond(client, { id, jsonrpc, result: { type: "stopped" } }); }).catch((cause) => { const error = getTRPCErrorFromUnknown(cause); opts.onError?.({ error, path, type, ctx, req, input }); respond(client, { id, jsonrpc, error: getErrorShape({ config: router._def._config, error, type, path, input, ctx }) }); abortController.abort(); }).finally(() => { clientAbortControllers.delete(id); }); respond(client, { id, jsonrpc, result: { type: "started" } }); } catch (cause) { const error = getTRPCErrorFromUnknown(cause); opts.onError?.({ error, path, type, ctx, req, input }); respond(client, { id, jsonrpc, error: getErrorShape({ config: router._def._config, error, type, path, input, ctx }) }); } } return { open(client) { client.data.abortController = new AbortController(); client.data.abortControllers = /* @__PURE__ */ new Map(); const url = createURL(client); client.data.ctx = createClientCtx.bind(null, client, url); const connectionParams = url.searchParams.get("connectionParams") === "1"; if (!connectionParams) { client.data.ctx = client.data.ctx(null); } }, async close(client) { client.data.abortController.abort(); await Promise.all( Array.from( client.data.abortControllers.values(), (ctrl) => ctrl.abort() ) ); }, async message(client, message) { const msgStr = message.toString(); if (msgStr === "PONG") { return; } if (msgStr === "PING") { client.send("PONG"); return; } try { const msgJSON = JSON.parse(msgStr); const msgs = Array.isArray(msgJSON) ? msgJSON : [msgJSON]; if (client.data.ctx instanceof Function) { const msg = msgs.shift(); client.data.ctx = client.data.ctx( parseConnectionParamsFromUnknown(msg.data) ); } const promises = []; for (const raw of msgs) { const msg = parseTRPCMessage( raw, router._def._config.transformer ); promises.push(handleRequest(client, msg)); } await Promise.all(promises); } catch (cause) { const error = new TRPCError({ code: "PARSE_ERROR", cause }); respond(client, { id: null, error: getErrorShape({ config: router._def._config, error, type: "unknown", path: void 0, input: void 0, ctx: void 0 }) }); } } }; } function isAsyncIterable(value) { return isObject(value) && Symbol.asyncIterator in value; } function run(fn) { return fn(); } function isObject(value) { return !!value && !Array.isArray(value) && typeof value === "object"; } function createURL(client) { try { const req = client.data.req; const protocol = client && "encrypted" in client && client.encrypted ? "https:" : "http:"; const host = req.headers.get("host") ?? "localhost"; return new URL(req.url, `${protocol}//${host}`); } catch (cause) { throw new TRPCError({ code: "BAD_REQUEST", message: "Invalid URL", cause }); } } // src/createBunServeHandler.ts function createBunServeHandler(opts, serveOptions) { const trpcHandler = createBunHttpHandler({ ...opts, emitWsUpgrades: true }); return { ...serveOptions, async fetch(req, server) { const trpcResponse = trpcHandler(req, server); if (trpcResponse) { return trpcResponse; } return serveOptions?.fetch?.call(server, req, server); }, websocket: createBunWSHandler(opts) }; } export { createBunHttpHandler, createBunServeHandler, createBunWSHandler }; //# sourceMappingURL=index.mjs.map