diff --git a/packages/coding-agent/src/modes/acp/acp-agent.ts b/packages/coding-agent/src/modes/acp/acp-agent.ts index fdbed3cc7..9d9572ff5 100644 --- a/packages/coding-agent/src/modes/acp/acp-agent.ts +++ b/packages/coding-agent/src/modes/acp/acp-agent.ts @@ -4,7 +4,9 @@ import { type AgentSideConnection, type AuthenticateRequest, type AuthenticateResponse, + type AuthMethod, type AvailableCommand, + type ClientCapabilities, type CloseSessionRequest, type CloseSessionResponse, type ForkSessionRequest, @@ -37,27 +39,34 @@ import { type SetSessionModeResponse, type Usage, } from "@agentclientprotocol/sdk"; -import type { Model } from "@oh-my-pi/pi-ai"; +import type { AssistantMessage, Model } from "@oh-my-pi/pi-ai"; import { logger, VERSION } from "@oh-my-pi/pi-utils"; import { disableProvider, enableProvider } from "../../capability"; import { Settings } from "../../config/settings"; import type { ExtensionUIContext } from "../../extensibility/extensions"; import { runExtensionCompact } from "../../extensibility/extensions/compact-handler"; +import { buildSkillPromptMessage, getSkillSlashCommandName } from "../../extensibility/skills"; import { loadSlashCommands } from "../../extensibility/slash-commands"; import { MCPManager } from "../../mcp/manager"; import type { MCPServerConfig } from "../../mcp/types"; import { loadAllExtensions } from "../../modes/components/extensions/state-manager"; import { theme } from "../../modes/theme/theme"; import type { AgentSession, AgentSessionEvent } from "../../session/agent-session"; +import { SKILL_PROMPT_MESSAGE_TYPE } from "../../session/messages"; import { SessionManager, type SessionInfo as StoredSessionInfo, type UsageStatistics, } from "../../session/session-manager"; +import { ACP_BUILTIN_SLASH_COMMANDS, executeAcpBuiltinSlashCommand } from "../../slash-commands/acp-builtins"; import { parseThinkingLevel } from "../../thinking"; +import { createAcpClientBridge } from "./acp-client-bridge"; import { mapAgentSessionEventToAcpSessionUpdates, mapToolKind } from "./acp-event-mapper"; +import { ACP_TERMINAL_AUTH_FLAG } from "./terminal-auth"; -const ACP_MODE_ID = "default"; +const ACP_DEFAULT_MODE_ID = "default"; +const ACP_PLAN_MODE_ID = "plan"; +const DEFAULT_PLAN_FILE_URL = "local://PLAN.md"; const MODE_CONFIG_ID = "mode"; const MODEL_CONFIG_ID = "model"; const THINKING_CONFIG_ID = "thinking"; @@ -84,7 +93,8 @@ type ManagedSessionRecord = { session: AgentSession; mcpManager: MCPManager | undefined; promptTurn: PromptTurnState | undefined; - liveMessageIds: WeakMap; + liveMessageId: string | undefined; + liveMessageProgress: { textEmitted: boolean; thoughtEmitted: boolean } | undefined; extensionsConfigured: boolean; }; @@ -152,6 +162,7 @@ export class AcpAgent implements Agent { #sessions = new Map(); #disposePromise: Promise | undefined; #cleanupRegistered = false; + #clientCapabilities: ClientCapabilities | undefined; constructor(connection: AgentSideConnection, initialSession: AgentSession, createSession: CreateAcpSession) { this.#connection = connection; @@ -159,8 +170,25 @@ export class AcpAgent implements Agent { this.#createSession = createSession; } - async initialize(_params: InitializeRequest): Promise { + async initialize(params: InitializeRequest): Promise { this.#registerConnectionCleanup(); + this.#clientCapabilities = params.clientCapabilities; + const authMethods: AuthMethod[] = [ + { + id: "agent", + name: "Use existing local credentials", + description: "Authenticate via the provider keys/OAuth state already configured under ~/.omp.", + }, + ]; + if (params.clientCapabilities?.auth?.terminal === true) { + authMethods.push({ + type: "terminal", + id: "terminal", + name: "Set up Oh My Pi in terminal", + description: "Launch the omp TUI to add provider keys and select models.", + args: [ACP_TERMINAL_AUTH_FLAG], + }); + } return { protocolVersion: PROTOCOL_VERSION, agentInfo: { @@ -168,13 +196,7 @@ export class AcpAgent implements Agent { title: "Oh My Pi", version: VERSION, }, - authMethods: [ - { - id: "agent", - name: "Agent-managed authentication", - description: "Oh My Pi uses its existing local authentication and provider configuration.", - }, - ], + authMethods, agentCapabilities: { loadSession: true, mcpCapabilities: { @@ -206,7 +228,7 @@ export class AcpAgent implements Agent { sessionId: record.session.sessionId, configOptions: this.#buildConfigOptions(record.session), models: this.#buildModelState(record.session), - modes: this.#buildModeState(), + modes: this.#buildModeState(record.session), }; this.#scheduleBootstrapUpdates(record.session.sessionId); return response; @@ -219,7 +241,7 @@ export class AcpAgent implements Agent { const response: LoadSessionResponse = { configOptions: this.#buildConfigOptions(record.session), models: this.#buildModelState(record.session), - modes: this.#buildModeState(), + modes: this.#buildModeState(record.session), }; this.#scheduleBootstrapUpdates(record.session.sessionId); return response; @@ -242,13 +264,13 @@ export class AcpAgent implements Agent { }; } - async unstable_resumeSession(params: ResumeSessionRequest): Promise { + async resumeSession(params: ResumeSessionRequest): Promise { this.#assertAbsoluteCwd(params.cwd); const record = await this.#resumeManagedSession(params.sessionId, params.cwd, params.mcpServers ?? []); const response: ResumeSessionResponse = { configOptions: this.#buildConfigOptions(record.session), models: this.#buildModelState(record.session), - modes: this.#buildModeState(), + modes: this.#buildModeState(record.session), }; this.#scheduleBootstrapUpdates(record.session.sessionId); return response; @@ -261,13 +283,13 @@ export class AcpAgent implements Agent { sessionId: record.session.sessionId, configOptions: this.#buildConfigOptions(record.session), models: this.#buildModelState(record.session), - modes: this.#buildModeState(), + modes: this.#buildModeState(record.session), }; this.#scheduleBootstrapUpdates(record.session.sessionId); return response; } - async unstable_closeSession(params: CloseSessionRequest): Promise { + async closeSession(params: CloseSessionRequest): Promise { const record = this.#sessions.get(params.sessionId); if (!record) { return {}; @@ -278,12 +300,17 @@ export class AcpAgent implements Agent { async setSessionMode(params: SetSessionModeRequest): Promise { const record = this.#getSessionRecord(params.sessionId); - if (params.modeId !== ACP_MODE_ID) { - throw new Error(`Unsupported ACP mode: ${params.modeId}`); - } + this.#applyModeChange(record.session, params.modeId); await this.#connection.sessionUpdate({ sessionId: record.session.sessionId, - update: this.#buildCurrentModeUpdate(), + update: this.#buildCurrentModeUpdate(record.session), + }); + await this.#connection.sessionUpdate({ + sessionId: record.session.sessionId, + update: { + sessionUpdate: "config_option_update", + configOptions: this.#buildConfigOptions(record.session), + }, }); return {}; } @@ -296,9 +323,7 @@ export class AcpAgent implements Agent { switch (params.configId) { case MODE_CONFIG_ID: - if (params.value !== ACP_MODE_ID) { - throw new Error(`Unsupported ACP mode config value: ${params.value}`); - } + this.#applyModeChange(record.session, params.value); break; case MODEL_CONFIG_ID: await this.#setModelById(record.session, params.value); @@ -356,13 +381,84 @@ export class AcpAgent implements Agent { void this.#handlePromptEvent(record, event); }); - record.session.prompt(converted.text, { images: converted.images }).catch((error: unknown) => { + this.#runPromptOrCommand(record, converted.text, converted.images).catch((error: unknown) => { this.#finishPrompt(record, undefined, error); }); return await pendingPrompt.promise; } + async #runPromptOrCommand(record: ManagedSessionRecord, text: string, images: AgentImageContent[]): Promise { + const skillResult = await this.#tryRunSkillCommand(record, text); + if (skillResult) { + return; + } + + const builtinResult = await executeAcpBuiltinSlashCommand(text, { + session: record.session, + sessionManager: record.session.sessionManager, + settings: Settings.instance, + cwd: record.session.sessionManager.getCwd(), + output: output => this.#emitCommandOutput(record, output), + refreshCommands: () => this.#emitAvailableCommandsUpdate(record), + notifyTitleChanged: async () => { + await this.#connection.sessionUpdate({ + sessionId: record.session.sessionId, + update: { + sessionUpdate: "session_info_update", + title: record.session.sessionName, + updatedAt: new Date().toISOString(), + }, + }); + }, + }); + if (builtinResult !== false) { + if ("prompt" in builtinResult) { + await record.session.prompt(builtinResult.prompt, { images }); + return; + } + const promptTurn = record.promptTurn; + this.#finishPrompt(record, { + stopReason: "end_turn", + usage: this.#buildTurnUsage( + promptTurn?.usageBaseline ?? + this.#cloneUsageStatistics(record.session.sessionManager.getUsageStatistics()), + record.session.sessionManager.getUsageStatistics(), + ), + userMessageId: promptTurn?.userMessageId, + }); + return; + } + + await record.session.prompt(text, { images }); + } + + async #tryRunSkillCommand(record: ManagedSessionRecord, text: string): Promise { + if (!text.startsWith("/skill:")) { + return false; + } + if (!record.session.skillsSettings?.enableSkillCommands) { + return false; + } + const spaceIndex = text.indexOf(" "); + const commandName = spaceIndex === -1 ? text.slice(1) : text.slice(1, spaceIndex); + const args = spaceIndex === -1 ? "" : text.slice(spaceIndex + 1).trim(); + const skillName = commandName.slice("skill:".length); + const skill = record.session.skills.find(candidate => candidate.name === skillName); + if (!skill) { + return false; + } + const built = await buildSkillPromptMessage(skill, args); + await record.session.promptCustomMessage({ + customType: SKILL_PROMPT_MESSAGE_TYPE, + content: built.message, + display: true, + details: built.details, + attribution: "user", + }); + return true; + } + async cancel(params: { sessionId: string }): Promise { const record = this.#getSessionRecord(params.sessionId); const promptTurn = record.promptTurn; @@ -384,7 +480,7 @@ export class AcpAgent implements Agent { async extMethod(method: string, params: { [key: string]: unknown }): Promise<{ [key: string]: unknown }> { switch (method) { - case "omp/sessions/listAll": { + case "_omp/sessions/listAll": { const limit = typeof params.limit === "number" ? Math.max(1, Math.min(5000, params.limit as number)) : 1000; const sessions = await SessionManager.listAll(); const sorted = sessions.sort((l, r) => r.modified.getTime() - l.modified.getTime()).slice(0, limit); @@ -393,7 +489,7 @@ export class AcpAgent implements Agent { total: sessions.length, }; } - case "omp/projects/list": { + case "_omp/projects/list": { const sessions = await SessionManager.listAll(); const buckets = new Map< string, @@ -421,7 +517,7 @@ export class AcpAgent implements Agent { const projects = Array.from(buckets.values()).sort((a, b) => b.lastActivityAt - a.lastActivityAt); return { projects, totalSessions: sessions.length }; } - case "omp/chats/byCwd": { + case "_omp/chats/byCwd": { const cwd = typeof params.cwd === "string" ? (params.cwd as string) : undefined; if (!cwd) throw new Error("cwd required"); const limit = typeof params.limit === "number" ? Math.max(1, Math.min(500, params.limit as number)) : 100; @@ -429,20 +525,20 @@ export class AcpAgent implements Agent { const sorted = sessions.sort((l, r) => r.modified.getTime() - l.modified.getTime()).slice(0, limit); return { sessions: sorted.map(s => this.#toSessionInfo(s)) }; } - case "omp/usage": { + case "_omp/usage": { const [firstRecord] = this.#sessions.values(); const target = firstRecord?.session ?? this.#initialSession; const reports = await target.fetchUsageReports(); return { reports: reports ?? [] }; } - case "omp/extensions": { + case "_omp/extensions": { const cwd = typeof params.cwd === "string" ? (params.cwd as string) : undefined; const sm = await Settings.init(); const disabledIds = (sm.get("disabledExtensions") as string[] | undefined) ?? []; const extensions = await loadAllExtensions(cwd, disabledIds); return { extensions: extensions as unknown as Array<{ [key: string]: unknown }> }; } - case "omp/extensions/toggle": { + case "_omp/extensions/toggle": { const providerId = params.providerId; if (typeof providerId !== "string") throw new Error("providerId required"); if (params.enabled === false) { @@ -562,6 +658,7 @@ export class AcpAgent implements Agent { async #registerPreparedSession(session: AgentSession, mcpServers: McpServer[]): Promise { const record = this.#createManagedSessionRecord(session); + session.setClientBridge(createAcpClientBridge(this.#connection, session.sessionId, this.#clientCapabilities)); try { await this.#configureExtensions(record); await this.#configureMcpServers(record, mcpServers); @@ -578,7 +675,8 @@ export class AcpAgent implements Agent { session, mcpManager: undefined, promptTurn: undefined, - liveMessageIds: new WeakMap(), + liveMessageId: undefined, + liveMessageProgress: undefined, extensionsConfigured: false, }; } @@ -627,33 +725,60 @@ export class AcpAgent implements Agent { return; } + this.#prepareLiveAssistantMessage(record, event); for (const notification of mapAgentSessionEventToAcpSessionUpdates(event, record.session.sessionId, { getMessageId: message => this.#getLiveMessageId(record, message), + getMessageProgress: message => this.#getLiveMessageProgress(record, message), })) { await this.#connection.sessionUpdate(notification); } + this.#clearLiveAssistantMessageAfterEvent(record, event); if (event.type === "agent_end") { await this.#emitEndOfTurnUpdates(record); this.#finishPrompt(record, { - stopReason: promptTurn.cancelRequested ? "cancelled" : "end_turn", + stopReason: this.#resolveStopReason(event, promptTurn.cancelRequested), usage: this.#buildTurnUsage(promptTurn.usageBaseline, record.session.sessionManager.getUsageStatistics()), userMessageId: promptTurn.userMessageId, }); } } + #prepareLiveAssistantMessage(record: ManagedSessionRecord, event: AgentSessionEvent): void { + if ( + (event.type === "message_start" || event.type === "message_update" || event.type === "message_end") && + event.message.role === "assistant" && + (event.type === "message_start" || !record.liveMessageId || !record.liveMessageProgress) + ) { + record.liveMessageId = crypto.randomUUID(); + record.liveMessageProgress = { textEmitted: false, thoughtEmitted: false }; + } + } + + #clearLiveAssistantMessageAfterEvent(record: ManagedSessionRecord, event: AgentSessionEvent): void { + if ((event.type === "message_end" && event.message.role === "assistant") || event.type === "agent_end") { + record.liveMessageId = undefined; + record.liveMessageProgress = undefined; + } + } + #getLiveMessageId(record: ManagedSessionRecord, message: unknown): string | undefined { if (typeof message !== "object" || message === null) { return undefined; } - const existing = record.liveMessageIds.get(message); - if (existing) { - return existing; + record.liveMessageId ??= crypto.randomUUID(); + return record.liveMessageId; + } + + #getLiveMessageProgress( + record: ManagedSessionRecord, + message: unknown, + ): { textEmitted: boolean; thoughtEmitted: boolean } | undefined { + if (typeof message !== "object" || message === null) { + return undefined; } - const nextMessageId = crypto.randomUUID(); - record.liveMessageIds.set(message, nextMessageId); - return nextMessageId; + record.liveMessageProgress ??= { textEmitted: false, thoughtEmitted: false }; + return record.liveMessageProgress; } #finishPrompt(record: ManagedSessionRecord, response?: PromptResponse, error?: unknown): void { @@ -671,6 +796,48 @@ export class AcpAgent implements Agent { promptTurn.resolve(response ?? { stopReason: "end_turn" }); } + #resolveStopReason( + event: Extract, + cancelRequested: boolean, + ): PromptResponse["stopReason"] { + if (cancelRequested) { + return "cancelled"; + } + const lastAssistant = [...event.messages] + .reverse() + .find((message): message is AssistantMessage => message.role === "assistant"); + const reason = lastAssistant?.stopReason; + switch (reason) { + case "aborted": + return "cancelled"; + case "length": + return "max_tokens"; + case "error": { + const errorMessage = lastAssistant?.errorMessage ?? ""; + if (/content[_ ]?filter|refus(al|ed)/i.test(errorMessage)) { + return "refusal"; + } + return "end_turn"; + } + default: + return "end_turn"; + } + } + + async #emitCommandOutput(record: ManagedSessionRecord, text: string): Promise { + if (!text) { + return; + } + await this.#connection.sessionUpdate({ + sessionId: record.session.sessionId, + update: { + sessionUpdate: "agent_message_chunk", + content: { type: "text", text }, + messageId: crypto.randomUUID(), + }, + }); + } + #assertAbsoluteCwd(cwd: string): void { if (!path.isAbsolute(cwd)) { throw new Error(`ACP cwd must be absolute: ${cwd}`); @@ -710,14 +877,20 @@ export class AcpAgent implements Agent { } #buildConfigOptions(session: AgentSession): SessionConfigOption[] { + const currentModeId = this.#getCurrentModeId(session); + const modeOptions = this.#getAvailableModes(session).map(mode => ({ + value: mode.id, + name: mode.name, + description: mode.description, + })); const configOptions: SessionConfigOption[] = [ { id: MODE_CONFIG_ID, name: "Mode", category: "mode", type: "select", - currentValue: ACP_MODE_ID, - options: [{ value: ACP_MODE_ID, name: "Default", description: "Standard ACP headless mode" }], + currentValue: currentModeId, + options: modeOptions, }, ]; @@ -805,17 +978,52 @@ export class AcpAgent implements Agent { return `${model.provider}/${model.id}`; } - #buildModeState(): SessionModeState { + #getAvailableModes(session: AgentSession): Array<{ id: string; name: string; description: string }> { + const modes = [{ id: ACP_DEFAULT_MODE_ID, name: "Default", description: "Standard ACP headless mode" }]; + if (Settings.instance.get("plan.enabled")) { + modes.push({ + id: ACP_PLAN_MODE_ID, + name: "Plan", + description: "Read-only planning mode that drafts a plan to a markdown file before any code changes", + }); + } + void session; + return modes; + } + + #getCurrentModeId(session: AgentSession): string { + return session.getPlanModeState()?.enabled ? ACP_PLAN_MODE_ID : ACP_DEFAULT_MODE_ID; + } + + #applyModeChange(session: AgentSession, modeId: string): void { + const availableModes = this.#getAvailableModes(session); + if (!availableModes.some(mode => mode.id === modeId)) { + throw new Error(`Unsupported ACP mode: ${modeId}`); + } + if (modeId === ACP_PLAN_MODE_ID) { + const previous = session.getPlanModeState(); + session.setPlanModeState({ + enabled: true, + planFilePath: previous?.planFilePath ?? DEFAULT_PLAN_FILE_URL, + workflow: previous?.workflow ?? "parallel", + reentry: previous !== undefined, + }); + } else { + session.setPlanModeState(undefined); + } + } + + #buildModeState(session: AgentSession): SessionModeState { return { - availableModes: [{ id: ACP_MODE_ID, name: "Default", description: "Standard ACP headless mode" }], - currentModeId: ACP_MODE_ID, + availableModes: this.#getAvailableModes(session), + currentModeId: this.#getCurrentModeId(session), }; } - #buildCurrentModeUpdate(): SessionUpdate { + #buildCurrentModeUpdate(session: AgentSession): SessionUpdate { return { sessionUpdate: "current_mode_update", - currentModeId: ACP_MODE_ID, + currentModeId: this.#getCurrentModeId(session), }; } @@ -838,6 +1046,20 @@ export class AcpAgent implements Agent { }); } + for (const command of ACP_BUILTIN_SLASH_COMMANDS) { + appendCommand(command); + } + + if (session.skillsSettings?.enableSkillCommands) { + for (const skill of session.skills) { + appendCommand({ + name: getSkillSlashCommandName(skill), + description: skill.description || `Run ${skill.name} skill`, + input: { hint: "arguments" }, + }); + } + } + for (const command of await loadSlashCommands({ cwd: session.sessionManager.getCwd() })) { appendCommand({ name: command.name, @@ -854,6 +1076,10 @@ export class AcpAgent implements Agent { cwd: session.cwd, title: session.title, updatedAt: session.modified.toISOString(), + _meta: { + messageCount: session.messageCount, + size: session.size, + }, }; } @@ -891,6 +1117,16 @@ export class AcpAgent implements Agent { }); } + async #emitAvailableCommandsUpdate(record: ManagedSessionRecord): Promise { + await this.#connection.sessionUpdate({ + sessionId: record.session.sessionId, + update: { + sessionUpdate: "available_commands_update", + availableCommands: await this.#buildAvailableCommands(record.session), + }, + }); + } + async #emitEndOfTurnUpdates(record: ManagedSessionRecord): Promise { const sessionId = record.session.sessionId;