diff --git a/packages/simulation/src/backend/control.ts b/packages/simulation/src/backend/control.ts index f5baa817be..a27a0cd2af 100644 --- a/packages/simulation/src/backend/control.ts +++ b/packages/simulation/src/backend/control.ts @@ -25,11 +25,11 @@ import { SimulationNetwork } from "./network" type ControlSocket = Bun.ServerWebSocket<{ unsubscribe?: () => void }> function parseRequest(input: string | Buffer) { - return SimulationProtocol.JsonRpc.decodeRequest(JSON.parse(typeof input === "string" ? input : input.toString())) + return SimulationProtocol.Backend.decodeRequest(JSON.parse(typeof input === "string" ? input : input.toString())) } -async function handle(socket: ControlSocket, request: SimulationProtocol.JsonRpc.Request): Promise { - switch (SimulationProtocol.Backend.decodeMethod(request.method)) { +async function handle(socket: ControlSocket, request: SimulationProtocol.Backend.Request): Promise { + switch (request.method) { case "llm.attach": { socket.data.unsubscribe?.() socket.data.unsubscribe = SimulationLLMExchange.subscribe((exchange) => { @@ -38,23 +38,22 @@ async function handle(socket: ControlSocket, request: SimulationProtocol.JsonRpc return { attached: true } } case "llm.chunk": { - const params = await SimulationProtocol.Backend.decodeChunkParams(request.params) await Effect.runPromise( SimulationLLMExchange.push( - params.id, - params.items.map((item) => ({ type: "item", item }) as const), + request.params.id, + request.params.items.map((item) => ({ type: "item", item }) as const), ), ) return { ok: true } } case "llm.finish": { - const params = await SimulationProtocol.Backend.decodeFinishParams(request.params) - await Effect.runPromise(SimulationLLMExchange.push(params.id, [{ type: "finish", reason: params.reason }])) + await Effect.runPromise( + SimulationLLMExchange.push(request.params.id, [{ type: "finish", reason: request.params.reason }]), + ) return { ok: true } } case "llm.disconnect": { - const params = await SimulationProtocol.Backend.decodeDisconnectParams(request.params) - await Effect.runPromise(SimulationLLMExchange.disconnect(params.id)) + await Effect.runPromise(SimulationLLMExchange.disconnect(request.params.id)) return { ok: true } } case "llm.pending": @@ -78,7 +77,7 @@ export function start(endpoint: string) { socket.data.unsubscribe?.() }, async message(socket, message) { - let request: SimulationProtocol.JsonRpc.Request | undefined + let request: SimulationProtocol.Backend.Request | undefined try { request = parseRequest(message) const result = await handle(socket, request) diff --git a/packages/simulation/src/frontend/server.ts b/packages/simulation/src/frontend/server.ts index 7b4baec95d..c717dbd4c6 100644 --- a/packages/simulation/src/frontend/server.ts +++ b/packages/simulation/src/frontend/server.ts @@ -7,16 +7,12 @@ export interface Server { readonly stop: () => void } -function actionParam(params: unknown) { - return SimulationProtocol.Frontend.decodeActionParams(params).action -} - function parseRequest(input: string | Buffer) { - return SimulationProtocol.JsonRpc.decodeRequest(JSON.parse(typeof input === "string" ? input : input.toString())) + return SimulationProtocol.Frontend.decodeRequest(JSON.parse(typeof input === "string" ? input : input.toString())) } -async function handle(harness: Harness, request: SimulationProtocol.JsonRpc.Request, headless: boolean) { - switch (SimulationProtocol.Frontend.decodeMethod(request.method)) { +async function handle(harness: Harness, request: SimulationProtocol.Frontend.Request, headless: boolean) { + switch (request.method) { case "ui.state": { if (headless) await harness.renderOnce() const result = SimulationActions.state(harness) @@ -24,7 +20,7 @@ async function handle(harness: Harness, request: SimulationProtocol.JsonRpc.Requ return result } case "ui.action": - return SimulationActions.execute(harness, actionParam(request.params)) + return SimulationActions.execute(harness, request.params.action) case "trace.list": return { records: SimulationTrace.list() } case "trace.clear": @@ -52,7 +48,7 @@ export function start(harness: Harness, endpoint: string, headless: boolean): Se SimulationTrace.add("control.disconnect") }, async message(socket, message) { - let request: SimulationProtocol.JsonRpc.Request | undefined + let request: SimulationProtocol.Frontend.Request | undefined try { request = parseRequest(message) const result = await handle(harness, request, headless) diff --git a/packages/simulation/src/protocol/index.ts b/packages/simulation/src/protocol/index.ts index b70e440e53..2534db8367 100644 --- a/packages/simulation/src/protocol/index.ts +++ b/packages/simulation/src/protocol/index.ts @@ -4,9 +4,12 @@ const JsonRpcID = Schema.Union([Schema.String, Schema.Number, Schema.Null]) type Json = Schema.Schema.Type export namespace JsonRpc { - export const Request = Schema.Struct({ + export const RequestFields = { jsonrpc: Schema.Literal("2.0"), id: Schema.optional(JsonRpcID), + } + export const Request = Schema.Struct({ + ...RequestFields, method: Schema.String, params: Schema.optional(Schema.Json), }) @@ -46,10 +49,6 @@ export namespace JsonRpc { } export namespace Frontend { - export const Method = Schema.Literals(["ui.state", "ui.action", "trace.list", "trace.clear", "trace.export"]) - export type Method = Schema.Schema.Type - export const decodeMethod = Schema.decodeUnknownSync(Method) - export const KeyModifiers = Schema.Struct({ ctrl: Schema.optional(Schema.Boolean), shift: Schema.optional(Schema.Boolean), @@ -96,7 +95,16 @@ export namespace Frontend { export const ActionParams = Schema.Struct({ action: Action }) export interface ActionParams extends Schema.Schema.Type {} - export const decodeActionParams = Schema.decodeUnknownSync(ActionParams) + + export const Request = Schema.Union([ + Schema.Struct({ ...JsonRpc.RequestFields, method: Schema.Literal("ui.action"), params: ActionParams }), + Schema.Struct({ + ...JsonRpc.RequestFields, + method: Schema.Literals(["ui.state", "trace.list", "trace.clear", "trace.export"]), + }), + ]) + export type Request = Schema.Schema.Type + export const decodeRequest = Schema.decodeUnknownSync(Request) export const TraceRecord = Schema.Struct({ id: Schema.Number, @@ -111,17 +119,6 @@ export namespace Frontend { } export namespace Backend { - export const Method = Schema.Literals([ - "llm.attach", - "llm.chunk", - "llm.finish", - "llm.disconnect", - "llm.pending", - "network.log", - ]) - export type Method = Schema.Schema.Type - export const decodeMethod = Schema.decodeUnknownSync(Method) - export const Item = Schema.Union([ Schema.Struct({ type: Schema.Literal("textDelta"), text: Schema.String }), Schema.Struct({ type: Schema.Literal("reasoningDelta"), text: Schema.String }), @@ -145,6 +142,18 @@ export namespace Backend { export const DisconnectParams = Schema.Struct({ id: Schema.String }) export interface DisconnectParams extends Schema.Schema.Type {} + export const Request = Schema.Union([ + Schema.Struct({ ...JsonRpc.RequestFields, method: Schema.Literal("llm.chunk"), params: ChunkParams }), + Schema.Struct({ ...JsonRpc.RequestFields, method: Schema.Literal("llm.finish"), params: FinishParams }), + Schema.Struct({ ...JsonRpc.RequestFields, method: Schema.Literal("llm.disconnect"), params: DisconnectParams }), + Schema.Struct({ + ...JsonRpc.RequestFields, + method: Schema.Literals(["llm.attach", "llm.pending", "network.log"]), + }), + ]) + export type Request = Schema.Schema.Type + export const decodeRequest = Schema.decodeUnknownSync(Request) + export const OpenedExchange = Schema.Struct({ id: Schema.String, url: Schema.String, body: Schema.Json }) export interface OpenedExchange extends Schema.Schema.Type {} @@ -156,9 +165,6 @@ export namespace Backend { }) export interface NetworkLogEntry extends Schema.Schema.Type {} - export const decodeChunkParams = Schema.decodeUnknownPromise(ChunkParams) - export const decodeFinishParams = Schema.decodeUnknownPromise(FinishParams) - export const decodeDisconnectParams = Schema.decodeUnknownPromise(DisconnectParams) } export * as SimulationProtocol from "./index"