feat(rpc): add validated desktop Pi RPC command palette
This commit is contained in:
@@ -20,6 +20,7 @@ const DEFAULT_EVENT_LIMIT = 1_000;
|
||||
const DEFAULT_CLOSE_TIMEOUT_MS = 2_000;
|
||||
const RESTORE_CONCURRENCY = 2;
|
||||
const RESOURCE_WARNING_COUNT = 6;
|
||||
const INITIAL_LAUNCH_GET_STATE_TIMEOUT_MS = 90_000;
|
||||
|
||||
export class UnknownAgentError extends Error {
|
||||
constructor(agentId) {
|
||||
@@ -229,6 +230,8 @@ function routeCommand(adapter, operation, payload = {}, runtimeState) {
|
||||
return adapter.send({ type: "set_thinking_level", level: payload.level });
|
||||
case "set_session_name":
|
||||
return adapter.send({ type: "set_session_name", name: payload.name });
|
||||
case "pi_rpc_command":
|
||||
return adapter.send({ type: payload.command, ...payload.input });
|
||||
case "extension_response":
|
||||
adapter.respondToExtension(payload.requestId, payload.response);
|
||||
return Promise.resolve({
|
||||
@@ -451,11 +454,16 @@ export function createAgentRegistry({
|
||||
let resolved;
|
||||
if (typeof nextPath === "string" && path.isAbsolute(nextPath)) {
|
||||
try {
|
||||
resolved = await ownedSessionPath(runtime.sessionDir, nextPath, {
|
||||
const candidate = await ownedSessionPath(runtime.sessionDir, nextPath, {
|
||||
allowMissing: true,
|
||||
});
|
||||
} catch {
|
||||
throw new Error("Pi reported a session outside its managed directory");
|
||||
await stat(candidate);
|
||||
resolved = candidate;
|
||||
} catch (error) {
|
||||
if (error?.code !== "ENOENT")
|
||||
throw new Error(
|
||||
"Pi reported a session outside its managed directory",
|
||||
);
|
||||
}
|
||||
}
|
||||
if (expectedSessionPath && resolved !== expectedSessionPath)
|
||||
@@ -627,8 +635,19 @@ export function createAgentRegistry({
|
||||
});
|
||||
}
|
||||
},
|
||||
onError: (error) => {
|
||||
onError: (error, { terminal = true } = {}) => {
|
||||
if (runtime.adapter !== adapter || runtime.stopped) return;
|
||||
if (!terminal) {
|
||||
const diagnostic = {
|
||||
code: error.code ?? "unknown",
|
||||
message: error.message,
|
||||
};
|
||||
publishAgent(runtime, "diagnostic", { error: diagnostic });
|
||||
publishWorkspace("runtime_diagnostic", runtime, {
|
||||
error: diagnostic,
|
||||
});
|
||||
return;
|
||||
}
|
||||
runtime.state = "error";
|
||||
runtime.stateVersion += 1;
|
||||
runtime.error = {
|
||||
@@ -673,7 +692,10 @@ export function createAgentRegistry({
|
||||
});
|
||||
runtime.adapter = adapter;
|
||||
try {
|
||||
const state = await adapter.send({ type: "get_state" });
|
||||
const state = await adapter.send(
|
||||
{ type: "get_state" },
|
||||
{ timeoutMs: INITIAL_LAUNCH_GET_STATE_TIMEOUT_MS },
|
||||
);
|
||||
runtime.state = state?.data?.isStreaming ? "streaming" : "idle";
|
||||
runtime.stateVersion += 1;
|
||||
await updateRuntimeIdentity(runtime, state, {
|
||||
@@ -1390,6 +1412,7 @@ export function createAgentRegistry({
|
||||
},
|
||||
async route(agentId, operation, payload = {}) {
|
||||
const runtime = getAgent(agentId);
|
||||
if (operation === "abort") return runtime.adapter.send({ type: "abort" });
|
||||
return enqueue(runtime, async () => {
|
||||
if (operation === "switch_session") {
|
||||
if (runtime.state === "streaming")
|
||||
@@ -1416,6 +1439,12 @@ export function createAgentRegistry({
|
||||
}
|
||||
const previousState = runtime.state;
|
||||
const previousStateVersion = runtime.stateVersion;
|
||||
const delivery =
|
||||
operation === "submit_prompt"
|
||||
? previousState === "streaming"
|
||||
? "follow_up"
|
||||
: "prompt"
|
||||
: undefined;
|
||||
const startsWork =
|
||||
operation === "prompt" ||
|
||||
(operation === "submit_prompt" && previousState !== "streaming");
|
||||
@@ -1447,7 +1476,7 @@ export function createAgentRegistry({
|
||||
runtime.attention =
|
||||
runtime.errorAttention || runtime.extensions.size > 0;
|
||||
}
|
||||
return response;
|
||||
return delivery ? { ...response, delivery } : response;
|
||||
} catch (error) {
|
||||
if (
|
||||
startsWork &&
|
||||
|
||||
@@ -2,28 +2,13 @@ import { randomUUID } from "node:crypto";
|
||||
import { spawn } from "node:child_process";
|
||||
import path from "node:path";
|
||||
import { StringDecoder } from "node:string_decoder";
|
||||
import { PI_RPC_COMMANDS } from "./pi-rpc-command-spec.js";
|
||||
|
||||
export const DEFAULT_COMMAND_TIMEOUT_MS = 30_000;
|
||||
export const COMPACT_COMMAND_TIMEOUT_MS = 5 * 60_000;
|
||||
export const MAX_PI_RPC_FRAME_BYTES = 1024 * 1024;
|
||||
|
||||
const supportedCommands = new Set([
|
||||
"prompt",
|
||||
"steer",
|
||||
"follow_up",
|
||||
"abort",
|
||||
"get_state",
|
||||
"get_session_stats",
|
||||
"new_session",
|
||||
"switch_session",
|
||||
"get_messages",
|
||||
"get_available_models",
|
||||
"get_commands",
|
||||
"set_session_name",
|
||||
"compact",
|
||||
"set_model",
|
||||
"set_thinking_level",
|
||||
]);
|
||||
const supportedCommands = new Set(PI_RPC_COMMANDS);
|
||||
|
||||
export class PiRpcError extends Error {
|
||||
constructor(code, message) {
|
||||
@@ -43,9 +28,9 @@ function assertAbsolutePath(value, field) {
|
||||
}
|
||||
}
|
||||
|
||||
function emitSafely(callback, value) {
|
||||
function emitSafely(callback, ...values) {
|
||||
try {
|
||||
callback(value);
|
||||
callback(...values);
|
||||
} catch {
|
||||
// Adapter observers must not interrupt the RPC reader.
|
||||
}
|
||||
@@ -161,6 +146,7 @@ export function startPiRpcAdapter({
|
||||
|
||||
let sequence = 0;
|
||||
let stdoutBuffer = "";
|
||||
let discardingOversizedFrame = false;
|
||||
let closed = false;
|
||||
let intentionalStop = false;
|
||||
const decoder = new StringDecoder("utf8");
|
||||
@@ -170,7 +156,7 @@ export function startPiRpcAdapter({
|
||||
resolveExit = resolve;
|
||||
});
|
||||
|
||||
const reportError = (error) => emitSafely(onError, error);
|
||||
const reportError = (error, metadata) => emitSafely(onError, error, metadata);
|
||||
const rejectPending = (error) => {
|
||||
for (const entry of pending.values()) {
|
||||
clearTimeout(entry.timeout);
|
||||
@@ -237,37 +223,58 @@ export function startPiRpcAdapter({
|
||||
if (frame.type === "response") handleResponse(frame);
|
||||
else handleEvent(frame);
|
||||
};
|
||||
|
||||
child.stdout.on("data", (chunk) => {
|
||||
if (closed) return;
|
||||
stdoutBuffer += decoder.write(chunk);
|
||||
const reportOversizedFrame = () =>
|
||||
reportError(
|
||||
new PiRpcError(
|
||||
"frame_too_large",
|
||||
`Pi RPC frame exceeds ${maxFrameBytes} bytes`,
|
||||
),
|
||||
{ terminal: false },
|
||||
);
|
||||
const consumeStdout = (text) => {
|
||||
let remaining = text;
|
||||
if (discardingOversizedFrame) {
|
||||
const newlineIndex = remaining.indexOf("\n");
|
||||
if (newlineIndex === -1) return;
|
||||
discardingOversizedFrame = false;
|
||||
remaining = remaining.slice(newlineIndex + 1);
|
||||
}
|
||||
stdoutBuffer += remaining;
|
||||
let newlineIndex;
|
||||
while ((newlineIndex = stdoutBuffer.indexOf("\n")) !== -1) {
|
||||
const line = stdoutBuffer.slice(0, newlineIndex);
|
||||
stdoutBuffer = stdoutBuffer.slice(newlineIndex + 1);
|
||||
if (Buffer.byteLength(line, "utf8") > maxFrameBytes) {
|
||||
reportOversizedFrame();
|
||||
continue;
|
||||
}
|
||||
handleLine(line);
|
||||
}
|
||||
if (Buffer.byteLength(stdoutBuffer, "utf8") > maxFrameBytes) {
|
||||
reportError(
|
||||
new PiRpcError(
|
||||
"frame_too_large",
|
||||
`Pi RPC frame exceeds ${maxFrameBytes} bytes`,
|
||||
),
|
||||
);
|
||||
reportOversizedFrame();
|
||||
stdoutBuffer = "";
|
||||
discardingOversizedFrame = true;
|
||||
}
|
||||
};
|
||||
|
||||
child.stdout.on("data", (chunk) => {
|
||||
if (!closed) consumeStdout(decoder.write(chunk));
|
||||
});
|
||||
child.stdout.on("end", () => {
|
||||
if (closed) return;
|
||||
const tail = stdoutBuffer + decoder.end();
|
||||
stdoutBuffer = "";
|
||||
if (tail.length > 0)
|
||||
consumeStdout(decoder.end());
|
||||
if (discardingOversizedFrame) {
|
||||
discardingOversizedFrame = false;
|
||||
return;
|
||||
}
|
||||
if (stdoutBuffer.length > 0)
|
||||
reportError(
|
||||
new PiRpcError(
|
||||
"unterminated_frame",
|
||||
"Pi RPC stdout ended without an LF-terminated frame",
|
||||
),
|
||||
);
|
||||
stdoutBuffer = "";
|
||||
});
|
||||
child.stderr.on("data", () => {});
|
||||
child.stderr.on("error", (error) =>
|
||||
@@ -293,7 +300,7 @@ export function startPiRpcAdapter({
|
||||
get sequence() {
|
||||
return sequence;
|
||||
},
|
||||
send(commandInput) {
|
||||
send(commandInput, { timeoutMs: requestedTimeoutMs } = {}) {
|
||||
if (!isRecord(commandInput) || typeof commandInput.type !== "string") {
|
||||
return Promise.reject(
|
||||
new PiRpcError(
|
||||
@@ -318,12 +325,24 @@ export function startPiRpcAdapter({
|
||||
),
|
||||
);
|
||||
}
|
||||
if (
|
||||
requestedTimeoutMs !== undefined &&
|
||||
(!Number.isSafeInteger(requestedTimeoutMs) || requestedTimeoutMs <= 0)
|
||||
) {
|
||||
return Promise.reject(
|
||||
new PiRpcError(
|
||||
"invalid_timeout",
|
||||
"Pi RPC command timeout must be a positive safe integer",
|
||||
),
|
||||
);
|
||||
}
|
||||
|
||||
const id = `bridge-${randomUUID()}`;
|
||||
const timeoutMs =
|
||||
commandInput.type === "compact"
|
||||
requestedTimeoutMs ??
|
||||
(commandInput.type === "compact"
|
||||
? compactCommandTimeoutMs
|
||||
: commandTimeoutMs;
|
||||
: commandTimeoutMs);
|
||||
return new Promise((resolve, reject) => {
|
||||
const timeout = setTimeout(() => {
|
||||
pending.delete(id);
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
const fields = {
|
||||
prompt: ["message", "images", "streamingBehavior"],
|
||||
steer: ["message", "images"],
|
||||
follow_up: ["message", "images"],
|
||||
abort: [],
|
||||
new_session: ["parentSession"],
|
||||
get_state: [],
|
||||
get_messages: [],
|
||||
set_model: ["provider", "modelId"],
|
||||
cycle_model: [],
|
||||
get_available_models: [],
|
||||
set_thinking_level: ["level"],
|
||||
cycle_thinking_level: [],
|
||||
set_steering_mode: ["mode"],
|
||||
set_follow_up_mode: ["mode"],
|
||||
compact: ["customInstructions"],
|
||||
set_auto_compaction: ["enabled"],
|
||||
set_auto_retry: ["enabled"],
|
||||
abort_retry: [],
|
||||
bash: ["command"],
|
||||
abort_bash: [],
|
||||
get_session_stats: [],
|
||||
export_html: ["outputPath"],
|
||||
switch_session: ["sessionPath"],
|
||||
fork: ["entryId"],
|
||||
clone: [],
|
||||
get_fork_messages: [],
|
||||
get_entries: ["since"],
|
||||
get_tree: [],
|
||||
get_last_assistant_text: [],
|
||||
set_session_name: ["name"],
|
||||
get_commands: [],
|
||||
};
|
||||
|
||||
export const PI_RPC_COMMANDS = Object.freeze(Object.keys(fields));
|
||||
|
||||
export function validatePiRpcCommand(command, input) {
|
||||
if (!Object.hasOwn(fields, command))
|
||||
throw new TypeError(`unsupported Pi RPC command: ${command}`);
|
||||
if (input === undefined) return {};
|
||||
if (input === null || typeof input !== "object" || Array.isArray(input))
|
||||
throw new TypeError("Pi RPC command input must be an object");
|
||||
for (const key of Object.keys(input)) {
|
||||
if (!fields[command].includes(key))
|
||||
throw new TypeError(`unsupported Pi RPC field: ${command}.${key}`);
|
||||
}
|
||||
return input;
|
||||
}
|
||||
@@ -1,4 +1,5 @@
|
||||
import path from "node:path";
|
||||
import { validatePiRpcCommand } from "../bridge/pi-rpc-command-spec.js";
|
||||
|
||||
export const PROTOCOL_VERSION = "v1";
|
||||
export const MAX_FRAME_BYTES = 64 * 1024;
|
||||
@@ -37,6 +38,7 @@ const requestOperations = new Map([
|
||||
["compact", { agent: true, payload: "compact" }],
|
||||
["set_model", { agent: true, payload: "model" }],
|
||||
["set_thinking_level", { agent: true, payload: "thinking" }],
|
||||
["pi_rpc_command", { agent: true, payload: "piRpcCommand" }],
|
||||
]);
|
||||
|
||||
const eventTypes = new Set([
|
||||
@@ -275,6 +277,17 @@ function validatePayload(kind, value) {
|
||||
);
|
||||
return { level };
|
||||
}
|
||||
case "piRpcCommand": {
|
||||
assertAllowedKeys(payload, new Set(["command", "input"]), "payload");
|
||||
const command = assertString(payload.command, "payload.command", {
|
||||
maxLength: 64,
|
||||
});
|
||||
try {
|
||||
return { command, input: validatePiRpcCommand(command, payload.input) };
|
||||
} catch (error) {
|
||||
throw new ProtocolError("invalid_message", error.message);
|
||||
}
|
||||
}
|
||||
default:
|
||||
throw new ProtocolError("invalid_message", "unsupported payload shape");
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user