返回 CodeWhale
protocol.ts
根目录 / crates / tui / extension-host / src / protocol.ts
1 /**
2 * Codewhale extension-host protocol, version 1.
3 *
4 * The Rust serde types in `crates/tui/src/extension_host/protocol.rs` are the
5 * source of truth. The constants, the method table, every params shape and the
6 * wire types come from `protocol.generated.ts`, which a Rust test renders from
7 * those types and fails on when the committed file drifts. Hand-written here:
8 * the frame codec, the JSON-RPC envelope checks, and the one rule Rust applies
9 * beyond its types (`host/hello`'s runtime name). Both sides also parse the
10 * shared corpus in `crates/tui/tests/fixtures/extension_host/protocol`.
11 *
12 * Frame: 4-byte magic `CWX1`, u32 little-endian payload length, UTF-8 JSON.
13 * Envelope: JSON-RPC 2.0.
14 */
15 import { MAGIC_ASCII, MAX_FRAME, HEADER_LEN, METHODS, SHAPES } from './protocol.generated.ts'
16 import type { Direction, HostTier, Kind, RpcErrorWire, Shape } from './protocol.generated.ts'
17
18 export * from './protocol.generated.ts'
19
20 export const MAGIC = Buffer.from(MAGIC_ASCII, 'ascii')
21
22 /** One row of the method table (`METHODS`). */
23 export interface MethodRow {
24 readonly name: string
25 readonly direction: Direction
26 readonly request: boolean
27 readonly params: string
28 readonly tiers: readonly HostTier[]
29 }
30
31 export type Message =
32 | { jsonrpc: '2.0'; id: number; method: string; params?: any }
33 | { jsonrpc: '2.0'; method: string; params?: any }
34 | { jsonrpc: '2.0'; id: number; result: any }
35 | { jsonrpc: '2.0'; id: number; error: RpcErrorWire }
36
37 export class FrameError extends Error {
38 constructor(message: string) {
39 super(message)
40 this.name = 'FrameError'
41 }
42 }
43
44 export class ProtocolError extends Error {
45 constructor(message: string) {
46 super(message)
47 this.name = 'ProtocolError'
48 }
49 }
50
51 /** Encode one message into a `CWX1` frame. Oversized payloads are refused, never truncated. */
52 export function encodeFrame(message: unknown): Buffer {
53 const payload = Buffer.from(JSON.stringify(message), 'utf8')
54 if (payload.length > MAX_FRAME) {
55 throw new FrameError(`frame of ${payload.length} bytes exceeds MAX_FRAME ${MAX_FRAME}`)
56 }
57 const header = Buffer.alloc(HEADER_LEN)
58 MAGIC.copy(header, 0)
59 header.writeUInt32LE(payload.length, 4)
60 return Buffer.concat([header, payload])
61 }
62
63 /** Incremental `CWX1` decoder. Throws `FrameError` on bad magic, bad length, or bad JSON. */
64 export class FrameDecoder {
65 private buffer: Buffer = Buffer.alloc(0)
66
67 push(chunk: Buffer): unknown[] {
68 this.buffer = this.buffer.length === 0 ? chunk : Buffer.concat([this.buffer, chunk])
69 const out: unknown[] = []
70 while (this.buffer.length >= HEADER_LEN) {
71 if (!this.buffer.subarray(0, 4).equals(MAGIC)) {
72 throw new FrameError('bad frame magic')
73 }
74 const length = this.buffer.readUInt32LE(4)
75 if (length > MAX_FRAME) throw new FrameError(`frame length ${length} exceeds MAX_FRAME`)
76 if (this.buffer.length < HEADER_LEN + length) break
77 const payload = this.buffer.subarray(HEADER_LEN, HEADER_LEN + length)
78 this.buffer = this.buffer.subarray(HEADER_LEN + length)
79 let value: unknown
80 try {
81 value = JSON.parse(payload.toString('utf8'))
82 } catch (error) {
83 throw new FrameError(`frame payload is not JSON: ${(error as Error).message}`)
84 }
85 out.push(value)
86 }
87 return out
88 }
89 }
90
91 // ---------------------------------------------------------------------------
92 // Validation, driven by the generated shapes. A shape is strict (unknown fields
93 // rejected) exactly where the Rust type has `deny_unknown_fields`: every
94 // host→core type and `OwnerRef`. The envelope is checked strictly for
95 // host→core messages and tolerantly for core→host ones.
96 // ---------------------------------------------------------------------------
97
98 function isObject(value: unknown): value is Record<string, unknown> {
99 return typeof value === 'object' && value !== null && !Array.isArray(value)
100 }
101
102 function checkShape(where: string, value: unknown, shape: Shape): asserts value is Record<string, any> {
103 if (!isObject(value)) throw new ProtocolError(`${where}: expected an object`)
104 for (const [key, kind] of Object.entries(shape.required)) {
105 if (!(key in value)) throw new ProtocolError(`${where}: missing field \`${key}\``)
106 checkKind(`${where}.${key}`, value[key], kind)
107 }
108 for (const [key, kind] of Object.entries(shape.optional)) {
109 // Every optional Rust field is an `Option` or a defaulted JSON value, so
110 // an explicit `null` reads as absent there too.
111 if (value[key] !== undefined && value[key] !== null) checkKind(`${where}.${key}`, value[key], kind)
112 }
113 if (shape.strict) {
114 for (const key of Object.keys(value)) {
115 if (!(key in shape.required) && !(key in shape.optional)) {
116 throw new ProtocolError(`${where}: unknown field \`${key}\``)
117 }
118 }
119 }
120 }
121
122 function checkKind(where: string, value: unknown, kind: Kind): void {
123 if (typeof kind === 'object') {
124 if ('ref' in kind) return checkShape(where, value, SHAPES[kind.ref])
125 if ('enum' in kind) {
126 if (typeof value !== 'string' || !kind.enum.includes(value)) {
127 throw new ProtocolError(`${where}: expected one of ${kind.enum.map((v) => `\`${v}\``).join(', ')}`)
128 }
129 return
130 }
131 if (!Array.isArray(value)) throw new ProtocolError(`${where}: expected an array`)
132 value.forEach((item, index) => checkKind(`${where}[${index}]`, item, kind.items))
133 return
134 }
135 switch (kind) {
136 case 'string':
137 if (typeof value !== 'string') throw new ProtocolError(`${where}: expected a string`)
138 return
139 case 'uint':
140 if (typeof value !== 'number' || !Number.isSafeInteger(value) || value < 0) {
141 throw new ProtocolError(`${where}: expected an unsigned integer`)
142 }
143 return
144 case 'integer':
145 if (typeof value !== 'number' || !Number.isSafeInteger(value)) throw new ProtocolError(`${where}: expected an integer`)
146 return
147 case 'boolean':
148 if (typeof value !== 'boolean') throw new ProtocolError(`${where}: expected a boolean`)
149 return
150 case 'object':
151 if (!isObject(value)) throw new ProtocolError(`${where}: expected an object`)
152 return
153 case 'json':
154 return
155 }
156 }
157
158 /**
159 * Validate one decoded message travelling in `direction` to or from a host of
160 * `tier`. Only methods in the generated table are admitted, each with its
161 * params shape and only on the tiers its row allows (a method reserved for the
162 * built-in tier is neither sent nor accepted by a plugin-tier host). Responses
163 * are validated as envelopes only: their result shape depends on the request,
164 * which the RPC layer checks. `methods` is the table to admit from; only a test
165 * of the tier rule passes anything but the generated one.
166 */
167 export function validateMessage(
168 value: unknown,
169 direction: Direction,
170 tier: HostTier,
171 methods: readonly MethodRow[] = METHODS,
172 ): Message {
173 const strict = direction === 'host_to_core'
174 if (!isObject(value)) throw new ProtocolError('message: expected an object')
175 if (value.jsonrpc !== '2.0') throw new ProtocolError('message: jsonrpc must be "2.0"')
176 const hasId = 'id' in value
177 if (hasId) checkKind('message.id', value.id, 'uint')
178 if ('method' in value) {
179 checkShape('message', value, { strict, required: { jsonrpc: 'string', method: 'string' }, optional: { id: 'uint', params: 'json' } })
180 const method = value.method as string
181 const spec = methods.find((entry) => entry.name === method && entry.direction === direction)
182 if (!spec) throw new ProtocolError(`unknown ${direction} method \`${method}\``)
183 if (!spec.tiers.includes(tier)) throw new ProtocolError(`\`${method}\` is not allowed on the ${tier} tier`)
184 if (spec.request !== hasId) {
185 throw new ProtocolError(`\`${method}\` must be ${spec.request ? 'a request (with id)' : 'a notification (no id)'}`)
186 }
187 const params = 'params' in value ? value.params : {}
188 checkShape(method, params, SHAPES[spec.params])
189 // Rust checks this after decoding `HelloParams` (`parse_host_message`).
190 if (method === 'host/hello' && params.runtime.name !== 'bun' && params.runtime.name !== 'node') {
191 throw new ProtocolError(`host/hello.runtime.name: unknown runtime \`${params.runtime.name}\``)
192 }
193 // Which spec fields each kind uses: `RegisterParams::check_spec`.
194 if (method === 'registry/register') {
195 const { kind, spec } = params
196 const reason =
197 kind === 'tool' && spec.input_schema == null
198 ? 'a tool registration needs `spec.input_schema`'
199 : kind === 'tool' && spec.argument_hint != null
200 ? 'a tool registration has no `spec.argument_hint`'
201 : kind === 'command' && spec.input_schema != null
202 ? 'a command registration has no `spec.input_schema`'
203 : (kind === 'hook' || kind === 'prompt_section' || kind === 'prompt_template' || kind === 'skill_root' || kind === 'shell_hook' || kind === 'mcp_server') && (spec.input_schema != null || spec.argument_hint != null)
204 ? 'a hook, prompt or skill root registration has no input schema or argument hint'
205 : undefined
206 if (reason !== undefined) throw new ProtocolError(`${method}: ${reason}`)
207 }
208 return value as Message
209 }
210 if (!hasId) throw new ProtocolError('response: missing id')
211 if ('error' in value) {
212 checkShape('message', value, { strict, required: { jsonrpc: 'string', id: 'uint', error: { ref: 'RpcErrorWire' } }, optional: {} })
213 return value as Message
214 }
215 if (!('result' in value)) throw new ProtocolError('response: needs result or error')
216 checkShape('message', value, { strict, required: { jsonrpc: 'string', id: 'uint', result: 'json' }, optional: {} })
217 return value as Message
218 }
219
219 lines TYPESCRIPT