Ver código fonte

refactor(simulation): attach RPC payload schemas

James Long 1 mês atrás
pai
commit
93529d2daa

+ 10 - 11
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<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
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)

+ 26 - 20
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<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"