| 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 |