refactor(simulation): attach RPC payload schemas

This commit is contained in:
James Long
2026-07-07 12:51:37 +00:00
parent adabb01caf
commit 93529d2daa
3 changed files with 41 additions and 40 deletions
+10 -11
View File
@@ -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<unknown> {
switch (SimulationProtocol.Backend.decodeMethod(request.method)) {
async function handle(socket: ControlSocket, request: SimulationProtocol.Backend.Request): Promise<unknown> {
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)
+5 -9
View File
@@ -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)
+26 -20
View File
@@ -4,9 +4,12 @@ const JsonRpcID = Schema.Union([Schema.String, Schema.Number, Schema.Null])
type Json = Schema.Schema.Type<typeof Schema.Json>
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<typeof Method>
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<typeof ActionParams> {}
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<typeof Request>
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<typeof Method>
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<typeof DisconnectParams> {}
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<typeof Request>
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<typeof OpenedExchange> {}
@@ -156,9 +165,6 @@ export namespace Backend {
})
export interface NetworkLogEntry extends Schema.Schema.Type<typeof NetworkLogEntry> {}
export const decodeChunkParams = Schema.decodeUnknownPromise(ChunkParams)
export const decodeFinishParams = Schema.decodeUnknownPromise(FinishParams)
export const decodeDisconnectParams = Schema.decodeUnknownPromise(DisconnectParams)
}
export * as SimulationProtocol from "./index"