diff --git a/src/bridge/agent-registry.js b/src/bridge/agent-registry.js index 4730317..539f32d 100644 --- a/src/bridge/agent-registry.js +++ b/src/bridge/agent-registry.js @@ -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 && diff --git a/src/bridge/pi-rpc-adapter.js b/src/bridge/pi-rpc-adapter.js index 428d7f3..f04280c 100644 --- a/src/bridge/pi-rpc-adapter.js +++ b/src/bridge/pi-rpc-adapter.js @@ -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); diff --git a/src/bridge/pi-rpc-command-spec.js b/src/bridge/pi-rpc-command-spec.js new file mode 100644 index 0000000..43d38cf --- /dev/null +++ b/src/bridge/pi-rpc-command-spec.js @@ -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; +} diff --git a/src/protocol/index.js b/src/protocol/index.js index 28a2413..3151bda 100644 --- a/src/protocol/index.js +++ b/src/protocol/index.js @@ -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"); } diff --git a/test/agent-registry.test.js b/test/agent-registry.test.js index ca9bd43..f9c5245 100644 --- a/test/agent-registry.test.js +++ b/test/agent-registry.test.js @@ -20,10 +20,12 @@ function createAdapterFactory() { startAdapter: (options) => { const adapter = { sent: [], + sentWithOptions: [], extensionResponses: [], stopped: false, - send(command) { + send(command, options) { this.sent.push(command); + this.sentWithOptions.push({ command, options }); return Promise.resolve({ type: "response", command: command.type, @@ -53,6 +55,10 @@ test("starts only the home agent and creates other worktree agents on explicit s }); const home = await registry.start(); + assert.deepEqual(fixture.calls[0].adapter.sentWithOptions[0], { + command: { type: "get_state" }, + options: { timeoutMs: 90_000 }, + }); assert.equal(registry.listAgents().length, 1); assert.equal(home.worktreePath, worktrees.home); @@ -72,6 +78,58 @@ test("starts only the home agent and creates other worktree agents on explicit s assert.ok(fixture.calls.every(({ adapter }) => adapter.stopped)); }); +test("sends abort without waiting behind a blocked prompt", async () => { + const worktrees = await createWorktrees(); + const fixture = createAdapterFactory(); + const registry = createAgentRegistry({ + homeWorktree: worktrees.home, + sessionRoot: worktrees.sessionRoot, + startAdapter: fixture.startAdapter, + }); + const agent = await registry.start(); + const adapter = fixture.calls[0].adapter; + const send = adapter.send.bind(adapter); + let releasePrompt; + adapter.send = (command) => { + if (command.type !== "prompt") return send(command); + adapter.sent.push(command); + return new Promise((resolve) => { + releasePrompt = () => + resolve({ type: "response", command: "prompt", success: true }); + }); + }; + const prompt = registry.route(agent.id, "prompt", { message: "Work" }); + while (!releasePrompt) await new Promise((resolve) => setImmediate(resolve)); + const abort = registry.route(agent.id, "abort"); + await new Promise((resolve) => setImmediate(resolve)); + assert.equal(adapter.sent.at(-1).type, "abort"); + releasePrompt(); + await Promise.all([prompt, abort]); + await registry.stop(); +}); + +test("keeps a runtime healthy after a non-terminal adapter diagnostic", async () => { + const worktrees = await createWorktrees(); + const fixture = createAdapterFactory(); + const registry = createAgentRegistry({ + homeWorktree: worktrees.home, + sessionRoot: worktrees.sessionRoot, + startAdapter: fixture.startAdapter, + }); + await registry.start(); + + fixture.calls[0].options.onError( + { code: "frame_too_large", message: "Pi RPC frame exceeds 32 bytes" }, + { terminal: false }, + ); + + const runtime = registry.getWorkspace().directories[0].runtimes[0]; + assert.equal(runtime.state, "idle"); + assert.equal(runtime.attention, false); + assert.equal(runtime.error, undefined); + await registry.stop(); +}); + test("stops agents created during shutdown and rejects new selections", async () => { const worktrees = await createWorktrees(); const fixture = createAdapterFactory(); @@ -270,8 +328,9 @@ test("coordinates forgetting with directory creation and active commands", async /runtimes are open/, ); const closing = registry.closeSessionRuntime( - registry.listAgents().find((agent) => agent.worktreePath === worktrees.feature) - .runtimeId, + registry + .listAgents() + .find((agent) => agent.worktreePath === worktrees.feature).runtimeId, ); await new Promise((resolve) => setImmediate(resolve)); releaseStop(); @@ -483,9 +542,10 @@ test("routes commands by explicit agent ID and replays only events after the cur type: "agent_state", data: { state: "streaming" }, }); - await registry.route(agent.id, "submit_prompt", { + const followUp = await registry.route(agent.id, "submit_prompt", { message: "After this turn", }); + assert.equal(followUp.delivery, "follow_up"); assert.deepEqual(adapter.sent.at(-1), { type: "follow_up", message: "After this turn", diff --git a/test/multi-session-registry.test.js b/test/multi-session-registry.test.js index 80653f3..316abf7 100644 --- a/test/multi-session-registry.test.js +++ b/test/multi-session-registry.test.js @@ -483,6 +483,45 @@ test("persists a first-prompt session when the runtime settles and awaits refres assert.match(manifestText, new RegExp(createdPath.replaceAll("/", "\\/"))); }); +test("does not persist an uncreated Pi session file across a restart", async () => { + const paths = await fixture(); + const futurePath = join( + sessionDirectoryPath(paths.sessionRoot, paths.home), + "future.jsonl", + ); + const adapters = adapterFactory({ + stateFor: () => ({ sessionFile: futurePath, sessionId: "future-id" }), + }); + const firstRegistry = createAgentRegistry({ + homeWorktree: paths.home, + sessionRoot: paths.sessionRoot, + startAdapter: adapters.startAdapter, + }); + await firstRegistry.start(); + assert.equal( + firstRegistry.getWorkspace().directories[0].runtimes[0].sessionPath, + undefined, + ); + await firstRegistry.stop(); + + const manifest = await readJson( + join(paths.sessionRoot, "bridge-workspace-v2.json"), + ); + assert.equal(manifest.runtimes[0].sessionPath, undefined); + + const restoredRegistry = createAgentRegistry({ + homeWorktree: paths.home, + sessionRoot: paths.sessionRoot, + startAdapter: adapters.startAdapter, + }); + await restoredRegistry.start(); + const restored = restoredRegistry.getWorkspace().directories[0].runtimes[0]; + assert.equal(restored.state, "idle"); + assert.equal(restored.sessionPath, undefined); + assert.equal(adapters.calls[1].options.sessionPath, undefined); + await restoredRegistry.stop(); +}); + test("extension responses do not clear unrelated error attention", async () => { const paths = await fixture(); const adapters = adapterFactory(); diff --git a/test/pi-rpc-adapter.test.js b/test/pi-rpc-adapter.test.js index f1c8f0d..7b0c73a 100644 --- a/test/pi-rpc-adapter.test.js +++ b/test/pi-rpc-adapter.test.js @@ -96,6 +96,32 @@ test("starts Pi in RPC mode and correlates a command response", async () => { await adapter.stop(); }); +test("allows a per-command timeout override without serializing it to Pi", async () => { + const fixture = createFixture(); + const adapter = startPiRpcAdapter({ + cwd: "/workspace/home", + sessionDir: "/workspace/sessions", + spawnProcess: fixture.spawn, + commandTimeoutMs: 10, + }); + const initial = adapter.send({ type: "get_state" }, { timeoutMs: 50 }); + const [initialCommand] = fixture.sent(); + await new Promise((resolve) => setTimeout(resolve, 20)); + fixture.child.stdout.write( + `${JSON.stringify({ type: "response", id: initialCommand.id, command: "get_state", success: true })}\n`, + ); + await initial; + assert.deepEqual(fixture.sent()[0], { + type: "get_state", + id: initialCommand.id, + }); + await assert.rejects( + adapter.send({ type: "get_session_stats" }), + /timed out/, + ); + await adapter.stop(); +}); + test("requests Pi session statistics for context and token status", async () => { const fixture = createFixture(); const adapter = startPiRpcAdapter({ @@ -182,6 +208,90 @@ test("forwards extension responses without replacing Pi's request ID", async () await adapter.stop(); }); +test("discards a fragmented oversized frame before processing the next frame", async () => { + const fixture = createFixture(); + const errors = []; + const events = []; + const adapter = startPiRpcAdapter({ + cwd: "/workspace/home", + sessionDir: "/workspace/sessions", + spawnProcess: fixture.spawn, + maxFrameBytes: 32, + onError: (error) => errors.push(error), + onEvent: (event) => events.push(event), + }); + const oversized = JSON.stringify({ + type: "message_update", + delta: "x".repeat(80), + }); + fixture.child.stdout.write(oversized.slice(0, 40)); + fixture.child.stdout.write( + `${oversized.slice(40)}\n${JSON.stringify({ type: "agent_settled" })}\n`, + ); + + assert.deepEqual( + errors.map((error) => error.code), + ["frame_too_large"], + ); + assert.deepEqual(events, [ + { + seq: 1, + type: "agent_state", + data: { state: "idle", event: { type: "agent_settled" } }, + }, + ]); + await adapter.stop(); +}); + +test("discards an oversized complete frame before processing the next frame", async () => { + const fixture = createFixture(); + const errors = []; + const events = []; + const adapter = startPiRpcAdapter({ + cwd: "/workspace/home", + sessionDir: "/workspace/sessions", + spawnProcess: fixture.spawn, + maxFrameBytes: 32, + onError: (error) => errors.push(error), + onEvent: (event) => events.push(event), + }); + const oversized = JSON.stringify({ + type: "message_update", + delta: "x".repeat(80), + }); + fixture.child.stdout.write( + `${oversized}\n${JSON.stringify({ type: "agent_settled" })}\n`, + ); + + assert.deepEqual( + errors.map((error) => error.code), + ["frame_too_large"], + ); + assert.equal(events.length, 1); + assert.equal(events[0].data.state, "idle"); + await adapter.stop(); +}); + +test("does not report an unterminated frame after discarding an oversized frame", async () => { + const fixture = createFixture(); + const errors = []; + const adapter = startPiRpcAdapter({ + cwd: "/workspace/home", + sessionDir: "/workspace/sessions", + spawnProcess: fixture.spawn, + maxFrameBytes: 32, + onError: (error) => errors.push(error), + }); + fixture.child.stdout.write("x".repeat(40)); + fixture.child.stdout.end(); + + assert.deepEqual( + errors.map((error) => error.code), + ["frame_too_large"], + ); + await adapter.stop(); +}); + test("reports invalid child output and rejects pending commands on child exit", async () => { const fixture = createFixture(); const errors = []; diff --git a/ui/src-tauri/src/bridge.rs b/ui/src-tauri/src/bridge.rs index 15328e2..ee372bf 100644 --- a/ui/src-tauri/src/bridge.rs +++ b/ui/src-tauri/src/bridge.rs @@ -375,7 +375,7 @@ pub async fn load_agent(socket_path: &str, agent_id: &str) -> Result Result<(), String> { +pub async fn submit_prompt(socket_path: &str, agent_id: &str, message: &str) -> Result { request( socket_path, "submit_prompt", @@ -383,7 +383,6 @@ pub async fn submit_prompt(socket_path: &str, agent_id: &str, message: &str) -> Some(json!({ "message": message })), ) .await - .map(|_| ()) } pub async fn abort(socket_path: &str, agent_id: &str) -> Result<(), String> { diff --git a/ui/src-tauri/src/lib.rs b/ui/src-tauri/src/lib.rs index 548687e..dbd7730 100644 --- a/ui/src-tauri/src/lib.rs +++ b/ui/src-tauri/src/lib.rs @@ -1,8 +1,8 @@ mod bridge; mod ui_state; -use serde_json::Value; -use std::sync::{Arc, Mutex}; +use serde_json::{json, Value}; +use std::{env, path::Path, sync::{Arc, Mutex}, time::Duration}; use tauri::{async_runtime::JoinHandle, AppHandle, Emitter, Manager, State}; struct LegacySubscription(Mutex>>); @@ -45,6 +45,58 @@ fn socket_path() -> Result { bridge::default_socket_path() } +fn requested_new_worktree_argument(args: &[String]) -> Option { + args.iter() + .position(|arg| arg == "--worktree") + .and_then(|index| args.get(index + 1)) + .cloned() +} + +fn requested_new_worktree(app: &AppHandle, args: &[String]) -> Option { + if !args.iter().any(|arg| arg == "--new") { + return None; + } + let saved_default = app + .path() + .app_data_dir() + .ok() + .and_then(|directory| ui_state::load_from(&directory).ok()) + .and_then(|state| state.default_new_session_worktree); + let worktree = requested_new_worktree_argument(args) + .or(saved_default) + .or_else(|| env::var("PI_STATUS_DEFAULT_WORKTREE").ok())?; + Path::new(&worktree).is_absolute().then_some(worktree) +} + +fn launch_new_session(app: AppHandle, args: &[String]) { + let Some(worktree_path) = requested_new_worktree(&app, args) else { return }; + tauri::async_runtime::spawn(async move { + // Give a newly-created webview time to register its frontend listeners. + tokio::time::sleep(Duration::from_millis(350)).await; + let _ = app.emit("workspace-new-session", json!({ + "phase": "starting", + "detail": "Starting a new Pi session…" + })); + match socket_path().map(|socket| (socket, worktree_path)) { + Ok((socket, worktree_path)) => match bridge::create_session_runtime(&socket, &worktree_path).await { + Ok(result) => { let _ = app.emit("workspace-new-session", json!({ + "phase": "ready", + "runtimeId": result.runtime.runtime_id, + "detail": "New session ready" + })); } + Err(error) => { let _ = app.emit("workspace-new-session", json!({ + "phase": "error", + "detail": format!("Could not start session: {error}") + })); } + }, + Err(error) => { let _ = app.emit("workspace-new-session", json!({ + "phase": "error", + "detail": error + })); } + } + }); +} + #[tauri::command] async fn get_workspace() -> Result { bridge::get_workspace(&socket_path()?).await @@ -146,7 +198,7 @@ async fn new_session(agent_id: String) -> Result { } #[tauri::command] -async fn submit_prompt(agent_id: String, message: String) -> Result<(), String> { +async fn submit_prompt(agent_id: String, message: String) -> Result { bridge::submit_prompt(&socket_path()?, &agent_id, &message).await } @@ -180,6 +232,16 @@ async fn set_session_name(agent_id: String, name: String) -> Result<(), String> bridge::set_session_name(&socket_path()?, &agent_id, &name).await } +#[tauri::command] +async fn pi_rpc_command(agent_id: String, command: String, input: Value) -> Result { + bridge::request( + &socket_path()?, + "pi_rpc_command", + Some(&agent_id), + Some(serde_json::json!({ "command": command, "input": input })), + ).await +} + #[tauri::command] async fn compact(agent_id: String, custom_instructions: Option) -> Result<(), String> { bridge::compact(&socket_path()?, &agent_id, custom_instructions.as_deref()).await @@ -281,6 +343,7 @@ fn unsubscribe_workspace(subscription: State<'_, WorkspaceSubscription>) -> Resu #[cfg_attr(mobile, tauri::mobile_entry_point)] pub fn run() { + let launch_args = env::args().collect::>(); let builder = tauri::Builder::default() .plugin(tauri_plugin_dialog::init()) .manage(LegacySubscription(Mutex::new(None))) @@ -289,6 +352,7 @@ pub fn run() { task: Mutex::new(None), }) .plugin(tauri_plugin_single_instance::init(|app, args, _cwd| { + launch_new_session(app.clone(), &args); if let Some(window) = app.get_webview_window("main") { match window_action(&args, window.is_visible().unwrap_or(false)) { WindowAction::Hide => { @@ -301,6 +365,10 @@ pub fn run() { } } })) + .setup(move |app| { + launch_new_session(app.handle().clone(), &launch_args); + Ok(()) + }) .invoke_handler(tauri::generate_handler![ get_workspace, get_workspace_summary, @@ -329,6 +397,7 @@ pub fn run() { set_thinking_level, set_session_name, compact, + pi_rpc_command, respond_to_extension, subscribe_agent ]); @@ -357,6 +426,14 @@ mod tests { assert_eq!(*tasks.lock().unwrap(), Some(2)); } + #[test] + fn new_session_uses_an_explicit_absolute_worktree() { + let args = vec!["pi-status-ui".into(), "--new".into(), "--worktree".into(), "/workspace".into()]; + assert_eq!(requested_new_worktree_argument(&args), Some("/workspace".into())); + let relative = vec!["pi-status-ui".into(), "--new".into(), "--worktree".into(), "workspace".into()]; + assert_eq!(requested_new_worktree_argument(&relative).filter(|path| Path::new(path).is_absolute()), None); + } + #[test] fn toggle_hides_a_visible_window_and_shows_a_hidden_window() { let toggle = vec!["--toggle".to_owned()]; diff --git a/ui/src-tauri/src/ui_state.rs b/ui/src-tauri/src/ui_state.rs index c3480b4..c9c8616 100644 --- a/ui/src-tauri/src/ui_state.rs +++ b/ui/src-tauri/src/ui_state.rs @@ -58,6 +58,8 @@ pub struct UiStateV1 { #[serde(default, skip_serializing_if = "Option::is_none")] pub interface_scale: Option, #[serde(default, skip_serializing_if = "Option::is_none")] + pub default_new_session_worktree: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] pub dismissed_collision_warning: Option, } @@ -73,6 +75,7 @@ impl Default for UiStateV1 { last_seen: BTreeMap::new(), workspace_cursor: None, interface_scale: None, + default_new_session_worktree: None, dismissed_collision_warning: None, } } @@ -127,6 +130,13 @@ pub fn validate(state: &UiStateV1) -> Result<(), String> { bounded(runtime_id, "selected runtime id")?; total_string_bytes = total_string_bytes.saturating_add(runtime_id.len()); } + if let Some(worktree) = &state.default_new_session_worktree { + bounded(worktree, "default new-session worktree")?; + if !Path::new(worktree).is_absolute() { + return Err("Default new-session worktree must be absolute".to_owned()); + } + total_string_bytes = total_string_bytes.saturating_add(worktree.len()); + } let mut total_draft_bytes = 0usize; for (runtime_id, draft) in &state.drafts { bounded(runtime_id, "draft runtime id")?; diff --git a/ui/src/App.css b/ui/src/App.css index 9f51ea9..2d82bc3 100644 --- a/ui/src/App.css +++ b/ui/src/App.css @@ -448,8 +448,18 @@ select:focus-visible { } .message.assistant { margin-right: 12px; - border-color: #505050; - background: #353535; + border-color: #496983; + background: #253544; +} +.message.tool, +.message.toolResult { + margin-right: 12px; + border-color: #2d8372; + background: #193936; +} +.message.error { + border-color: #bf6570; + background: #44282d; } .message strong { display: block; @@ -461,6 +471,16 @@ select:focus-visible { .message.user strong { color: #ffc170; } +.message.assistant strong { + color: #b9d9ff; +} +.message.tool strong, +.message.toolResult strong { + color: #a9f2df; +} +.message.error strong { + color: #ffd7db; +} .message pre { margin: 0; white-space: pre-wrap; @@ -469,6 +489,30 @@ select:focus-visible { font-size: 13px; line-height: 1.32; } +.tool-result summary { + display: flex; + align-items: center; + justify-content: space-between; + cursor: pointer; + list-style: none; + color: #a9f2df; + font-size: 10px; + font-weight: 700; + text-transform: uppercase; +} +.tool-result summary::-webkit-details-marker { + display: none; +} +.tool-result summary::after { + content: "▾"; + font-size: 14px; +} +.tool-result[open] summary::after { + content: "▴"; +} +.tool-result[open] summary { + margin-bottom: 6px; +} .pending-message { border-style: dashed; } @@ -529,6 +573,34 @@ select:focus-visible { align-items: center; font-size: 10px; } +.pi-controls select { + min-width: 0; + appearance: none; + color-scheme: dark; + padding: 5px 25px 5px 6px; + border: 1px solid #70522f; + border-radius: 4px; + background: #2c2a27 + url("data:image/svg+xml,%3Csvg xmlns='http://www.w3.org/2000/svg' width='12' height='8' viewBox='0 0 12 8'%3E%3Cpath d='m1 1 5 5 5-5' fill='none' stroke='%23f0a347' stroke-linecap='round' stroke-linejoin='round' stroke-width='2'/%3E%3C/svg%3E") + no-repeat right 7px center; + color: #f2ece4; +} +.pi-controls select:hover { + border-color: #d88735; + background-color: #383129; +} +.pi-controls select:focus { + border-color: #f0a347; + box-shadow: 0 0 0 2px rgba(240, 163, 71, 0.22); +} +.pi-controls select option { + background: #2c2a27; + color: #f2ece4; +} +.pi-controls select option:checked { + background: #4b3420; + color: #fff2df; +} .composer-progress { display: flex; align-items: center; @@ -582,12 +654,94 @@ select:focus-visible { display: flex; gap: 7px; } +.default-worktree-control { + display: grid; + gap: 4px; + margin: 8px 0; + color: #ddd8d0; + font-size: 11px; + font-weight: 700; +} +.default-worktree-control select { + min-width: 0; + appearance: none; + color-scheme: dark; + padding: 5px 28px 5px 7px; + border: 1px solid #70522f; + border-radius: 4px; + background: #2c2a27 + url("data:image/svg+xml,%3Csvg xmlns='http://www.w3.org/2000/svg' width='12' height='8' viewBox='0 0 12 8'%3E%3Cpath d='m1 1 5 5 5-5' fill='none' stroke='%23f0a347' stroke-linecap='round' stroke-linejoin='round' stroke-width='2'/%3E%3C/svg%3E") + no-repeat right 8px center; + color: #f2ece4; +} +.default-worktree-control select:hover { + border-color: #d88735; + background-color: #383129; +} +.default-worktree-control select:focus { + border-color: #f0a347; + box-shadow: 0 0 0 2px rgba(240, 163, 71, 0.22); +} +.default-worktree-control option { + background: #2c2a27; + color: #f2ece4; +} +.default-worktree-control option:checked { + background: #4b3420; + color: #fff2df; +} +.default-worktree-control small { + color: #b9afa4; + font-weight: 400; +} .extension-widget { padding: 5px 8px; border-left: 3px solid #f0a347; background: #302a24; font-size: 12px; } +.startup-blocker { + position: absolute; + z-index: 20; + inset: 0; + display: grid; + place-items: center; + padding: 20px; + background: rgba(20, 18, 16, 0.82); + backdrop-filter: blur(3px); + outline: 0; +} +.startup-blocker-card { + display: flex; + align-items: center; + gap: 12px; + width: min(380px, 90vw); + padding: 16px; + border: 1px solid #f0a347; + border-radius: 8px; + background: #33271d; + box-shadow: 0 12px 36px rgba(0, 0, 0, 0.45); + color: #ffd9aa; +} +.startup-blocker-card > div { + display: grid; + gap: 3px; +} +.startup-blocker-card strong { + color: #fff2df; + font-size: 14px; +} +.startup-blocker-card p { + margin: 0; + font-weight: 700; +} +.startup-blocker-card small { + color: #d8b892; +} +.startup-bars { + flex: 0 0 auto; + transform: scale(1.35); +} .modal-backdrop, .extension-backdrop { position: absolute; @@ -611,6 +765,9 @@ select:focus-visible { border-radius: 10px; background: #2d2e31; } +.session-picker { + overflow: hidden; +} .picker-heading { display: flex; justify-content: space-between; @@ -619,20 +776,105 @@ select:focus-visible { .extension h2 { margin: 0; } +.session-search input { + box-sizing: border-box; + width: 100%; + padding: 8px 10px; + border: 1px solid #5f6570; + border-radius: 6px; + background: #202125; + color: #f1f3f4; +} +.session-result-count { + margin: 0; + color: #8d97a6; + font-size: 0.85rem; +} .session-list { display: grid; - gap: 5px; + gap: 6px; + max-height: min(44vh, 430px); + margin: 0; + padding: 0; + overflow-y: auto; + list-style: none; } -.session-list button { +.session-row { display: flex; + align-items: center; justify-content: space-between; + width: 100%; + gap: 12px; + padding: 10px; + border: 1px solid #474b54; + border-radius: 7px; + background: #25262a; text-align: left; } -.session-list small { +.session-row:hover { + border-color: #b76f26; + background: #303136; +} +.session-row-content { + display: grid; + min-width: 0; + gap: 3px; +} +.session-preview, +.session-metadata { + overflow: hidden; + color: #a7afbc; + text-overflow: ellipsis; + white-space: nowrap; +} +.session-metadata { + display: flex; + gap: 8px; color: #8d97a6; } +.session-open { + flex: 0 0 auto; + padding: 2px 6px; + border-radius: 999px; + background: #234a33; + color: #a5e3bc; + font-size: 0.75rem; + font-weight: 700; +} +.extension-eyebrow, +.extension-request-label, +.extension-request-action, +.extension-request-details { + margin: 0; +} +.extension-eyebrow, +.extension-request-label { + color: #b9afa4; + font-size: 0.75rem; + font-weight: 800; + letter-spacing: 0.08em; + text-transform: uppercase; +} +.extension-request { + display: grid; + gap: 6px; + padding: 12px; + border-left: 3px solid #f0a347; + border-radius: 4px; + background: #25262a; +} +.extension-request-action { + color: #fff3df; + font-size: 1.05rem; + font-weight: 750; +} +.extension-request-details { + white-space: pre-wrap; + overflow-wrap: anywhere; +} .extension-options { display: flex; + flex-wrap: wrap; gap: 7px; } .extension textarea { diff --git a/ui/src/App.tsx b/ui/src/App.tsx index 3df0976..e68fa76 100644 --- a/ui/src/App.tsx +++ b/ui/src/App.tsx @@ -6,6 +6,8 @@ import "./App.css"; import { ConversationWorkspace } from "./components/ConversationWorkspace"; import { DirectorySidebar } from "./components/DirectorySidebar"; import { ExtensionDialog } from "./components/ExtensionDialog"; +import { CommandFormDialog } from "./components/CommandFormDialog"; +import type { RpcCommand } from "./commands/rpc"; import { SessionPicker } from "./components/SessionPicker"; import { SessionTabs } from "./components/SessionTabs"; import type { @@ -37,6 +39,7 @@ type LocalOperation = { export default function App() { const { state, dispatch, refresh, loadSnapshot } = useWorkspace(); const [status, setStatus] = useState("Connecting to Pi Status Bridge…"); + const [commandForm, setCommandForm] = useState(); const [view, setView] = useState<"conversation" | "settings">("conversation"); const [adding, setAdding] = useState(false); const [folderPath, setFolderPath] = useState(""); @@ -64,6 +67,7 @@ export default function App() { const mountedRef = useRef(true); const operationSequence = useRef(0); const [operation, setOperation] = useState(); + const [abortingRuntimeId, setAbortingRuntimeId] = useState(); const [restoreTabFocus, setRestoreTabFocus] = useState(false); const startOperation = ( kind: LocalOperation["kind"], @@ -99,6 +103,14 @@ export default function App() { ); }, [selected, dispatch, operation]); + useEffect(() => { + if ( + abortingRuntimeId && + state.runtimesById[abortingRuntimeId]?.summary.state !== "streaming" + ) + setAbortingRuntimeId(undefined); + }, [abortingRuntimeId, state.runtimesById]); + async function createRuntime(path = state.selectedDirectoryPath) { if (!path || !mountedRef.current) return; const operationId = startOperation( @@ -292,9 +304,13 @@ export default function App() { } const runtimeId = selected.summary.runtimeId, id = ++submissionId, + initialDelivery = + selected.summary.state === "streaming" ? "follow_up" : "prompt", operationId = startOperation( "submitting", - "Sending prompt to Pi…", + initialDelivery === "follow_up" + ? "Queueing follow-up…" + : "Sending prompt to Pi…", selected.summary.runtimeId, ); dispatch({ @@ -307,24 +323,49 @@ export default function App() { (message) => message.role === "user", ).length, phase: "sending", + delivery: initialDelivery, }, }); dispatch({ type: "draftChanged", runtimeId, draft: "" }); - const sent = await agentCommand( - "submit_prompt", - { message: text }, - false, - selected, - false, - ); - if (sent) { - dispatch({ type: "submissionSent", runtimeId, id }); - setStatus("Prompt accepted · waiting for Pi…"); - } else { + try { + const result = await invoke<{ delivery?: "prompt" | "follow_up" }>( + "submit_prompt", + { agentId: selected.summary.agentId, message: text }, + ); + const delivery = result.delivery ?? initialDelivery; + dispatch({ type: "submissionSent", runtimeId, id, delivery }); + setStatus( + delivery === "follow_up" + ? "Follow-up queued · waiting for Pi…" + : "Prompt accepted · waiting for Pi…", + ); + } catch (error) { + setStatus(`Pi command failed: ${String(error)}`); dispatch({ type: "submissionRemoved", runtimeId, id }); dispatch({ type: "draftChanged", runtimeId, draft: text }); + } finally { + finishOperation(operationId); } - finishOperation(operationId); + } + async function abortSelected() { + const target = selected; + if (!target?.summary.agentId) return; + setAbortingRuntimeId(target.summary.runtimeId); + setStatus("Aborting Pi…"); + try { + await invoke("abort", { agentId: target.summary.agentId }); + } catch (error) { + setStatus(`Pi command failed: ${String(error)}`); + setAbortingRuntimeId(undefined); + } + } + async function runRpcCommand( + command: RpcCommand, + input: Record = {}, + ) { + if (!selected) return; + await agentCommand("pi_rpc_command", { command: command.command, input }); + setCommandForm(undefined); } async function respond(response: Record) { const target = selected; @@ -342,10 +383,15 @@ export default function App() { const pendingSubmission = pendingSubmissions[pendingSubmissions.length - 1]; const displayedStatus = operation?.detail ?? + state.newSessionLaunch?.detail ?? (pendingSubmission ? pendingSubmission.phase === "sending" - ? "Sending prompt to Pi…" - : "Prompt accepted · waiting for Pi…" + ? pendingSubmission.delivery === "follow_up" + ? "Queueing follow-up…" + : "Sending prompt to Pi…" + : pendingSubmission.delivery === "follow_up" + ? "Follow-up queued · waiting for Pi…" + : "Prompt accepted · waiting for Pi…" : undefined) ?? (selected?.progress.phase === "working" || selected?.progress.phase === "recovering" @@ -353,11 +399,19 @@ export default function App() { : undefined) ?? status; const isWorking = + state.newSessionLaunch?.phase === "starting" || selected?.progress.phase === "working" || selected?.progress.phase === "recovering"; + const isStartingSession = + state.newSessionLaunch?.phase === "starting" || + operation?.kind === "creating"; const sameDirectoryCollision = directory && directory.openCount > 1; return ( -
+
{ @@ -379,15 +433,21 @@ export default function App() { {isWorking && ( )}
{isWorking && ( -
); } diff --git a/ui/src/commands/rpc.ts b/ui/src/commands/rpc.ts new file mode 100644 index 0000000..139d932 --- /dev/null +++ b/ui/src/commands/rpc.ts @@ -0,0 +1,141 @@ +export type RpcField = { + name: string; + label: string; + kind: "text" | "path" | "boolean" | "enum" | "json"; + required?: boolean; + options?: string[]; +}; + +export type RpcCommand = { + command: string; + category: + | "Prompting" + | "Session" + | "Model & Thinking" + | "Queue" + | "Context & Retry" + | "Shell" + | "Inspection & Export"; + description: string; + input?: RpcField[]; +}; + +const text = (name: string, required = true): RpcField => ({ + name, + label: name, + kind: "text", + required, +}); +const path = (name: string, required = true): RpcField => ({ + name, + label: name, + kind: "path", + required, +}); +const boolean = (name: string): RpcField => ({ + name, + label: name, + kind: "boolean", + required: true, +}); +const enumField = (name: string, options: string[]): RpcField => ({ + name, + label: name, + kind: "enum", + options, + required: true, +}); +const json = (name: string): RpcField => ({ name, label: name, kind: "json" }); + +const define = ( + category: RpcCommand["category"], + commands: Array<[string, string, RpcField[]?]>, +) => + commands.map(([command, description, input]) => ({ + command, + category, + description, + ...(input ? { input } : {}), + })); + +export const rpcCommandCatalog: RpcCommand[] = [ + ...define("Prompting", [ + [ + "prompt", + "Send a prompt", + [ + text("message"), + json("images"), + enumField("streamingBehavior", ["steer", "followUp"]), + ], + ], + ["steer", "Queue a steering message", [text("message"), json("images")]], + [ + "follow_up", + "Queue a follow-up message", + [text("message"), json("images")], + ], + ["abort", "Abort the agent"], + ]), + ...define("Session", [ + ["new_session", "Start a new session", [path("parentSession", false)]], + ["switch_session", "Switch session", [path("sessionPath")]], + ["fork", "Fork from a message", [text("entryId")]], + ["clone", "Clone the active branch"], + ["set_session_name", "Set session name", [text("name")]], + ["get_fork_messages", "List forkable messages"], + ["get_entries", "Read session entries", [text("since", false)]], + ["get_tree", "Read session tree"], + ["get_state", "Read session state"], + ["get_messages", "Read messages"], + ["get_last_assistant_text", "Read last assistant text"], + ]), + ...define("Model & Thinking", [ + ["get_available_models", "List models"], + ["set_model", "Select model", [text("provider"), text("modelId")]], + ["cycle_model", "Cycle model"], + [ + "set_thinking_level", + "Set thinking level", + [ + enumField("level", [ + "off", + "minimal", + "low", + "medium", + "high", + "xhigh", + "max", + ]), + ], + ], + ["cycle_thinking_level", "Cycle thinking level"], + ]), + ...define("Queue", [ + [ + "set_steering_mode", + "Set steering delivery", + [enumField("mode", ["all", "one-at-a-time"])], + ], + [ + "set_follow_up_mode", + "Set follow-up delivery", + [enumField("mode", ["all", "one-at-a-time"])], + ], + ]), + ...define("Context & Retry", [ + ["compact", "Compact context", [text("customInstructions", false)]], + ["set_auto_compaction", "Set automatic compaction", [boolean("enabled")]], + ["set_auto_retry", "Set automatic retry", [boolean("enabled")]], + ["abort_retry", "Abort retry"], + ]), + ...define("Shell", [ + ["bash", "Run shell command", [text("command")]], + ["abort_bash", "Abort shell command"], + ]), + ...define("Inspection & Export", [ + ["get_session_stats", "Read session statistics"], + ["export_html", "Export session HTML", [path("outputPath", false)]], + ["get_commands", "List Pi commands"], + ]), +]; diff --git a/ui/src/components/CommandFormDialog.tsx b/ui/src/components/CommandFormDialog.tsx new file mode 100644 index 0000000..0de79aa --- /dev/null +++ b/ui/src/components/CommandFormDialog.tsx @@ -0,0 +1,160 @@ +import { useEffect, useRef, useState } from "react"; +import type { RpcCommand, RpcField } from "../commands/rpc"; + +const controls = (container: HTMLElement | null) => + Array.from( + container?.querySelectorAll( + "button, input, select, textarea", + ) ?? [], + ).filter((node) => !node.hasAttribute("disabled")); + +function valueFor(field: RpcField, value: string | boolean): unknown { + if (field.kind === "boolean") return value === true; + if (field.kind === "json") + return value ? JSON.parse(value as string) : undefined; + return value || undefined; +} + +export function CommandFormDialog({ + command, + onSubmit, + onCancel, +}: { + command?: RpcCommand; + onSubmit: (input: Record) => void; + onCancel: () => void; +}) { + const ref = useRef(null); + const opener = useRef(null); + const [values, setValues] = useState>({}); + const [error, setError] = useState(); + useEffect(() => { + if (!command) return; + opener.current = document.activeElement as HTMLElement | null; + setValues({}); + setError(undefined); + const frame = requestAnimationFrame(() => + controls(ref.current)[0]?.focus(), + ); + return () => { + cancelAnimationFrame(frame); + opener.current?.focus(); + }; + }, [command]); + if (!command) return null; + const fields = command.input ?? []; + const submit = () => { + try { + const input: Record = {}; + for (const field of fields) { + const value = valueFor( + field, + values[field.name] ?? (field.kind === "boolean" ? false : ""), + ); + if (field.required && (value === undefined || value === "")) + throw new Error(`${field.label} is required`); + if (value !== undefined) input[field.name] = value; + } + onSubmit(input); + } catch (cause) { + setError(cause instanceof Error ? cause.message : "Invalid input"); + } + }; + return ( +
+
{ + if (event.key === "Escape") { + event.preventDefault(); + onCancel(); + } + if (event.key === "Tab") { + const items = controls(ref.current), + first = items[0], + last = items[items.length - 1]; + if (!first || !last) return; + if (event.shiftKey && document.activeElement === first) { + event.preventDefault(); + last.focus(); + } else if (!event.shiftKey && document.activeElement === last) { + event.preventDefault(); + first.focus(); + } + } + }} + > +

Pi RPC command

+

/{command.command}

+

{command.description}

+ {fields.map((field) => ( +