feat(acp): extend AcpAgent with session management and model negotiation

- Adds create/resume/list/page session lifecycle handling to AcpAgent
- Implements mode switching (default vs. plan), MCP server configuration, and model/thinking config negotiation with the connected client
- Routes incoming ACP connections to AgentSession via the new ClientBridge
This commit is contained in:
Ogrodev
2026-05-11 12:31:52 -03:00
committed by can1357
parent b707cdaa08
commit f81d9b7fef
+282 -46
View File
@@ -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<object, string>;
liveMessageId: string | undefined;
liveMessageProgress: { textEmitted: boolean; thoughtEmitted: boolean } | undefined;
extensionsConfigured: boolean;
};
@@ -152,6 +162,7 @@ export class AcpAgent implements Agent {
#sessions = new Map<string, ManagedSessionRecord>();
#disposePromise: Promise<void> | 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<InitializeResponse> {
async initialize(params: InitializeRequest): Promise<InitializeResponse> {
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<ResumeSessionResponse> {
async resumeSession(params: ResumeSessionRequest): Promise<ResumeSessionResponse> {
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<CloseSessionResponse> {
async closeSession(params: CloseSessionRequest): Promise<CloseSessionResponse> {
const record = this.#sessions.get(params.sessionId);
if (!record) {
return {};
@@ -278,12 +300,17 @@ export class AcpAgent implements Agent {
async setSessionMode(params: SetSessionModeRequest): Promise<SetSessionModeResponse> {
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<void> {
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<boolean> {
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<void> {
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<ManagedSessionRecord> {
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<object, string>(),
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<AgentSessionEvent, { type: "agent_end" }>,
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<void> {
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<void> {
await this.#connection.sessionUpdate({
sessionId: record.session.sessionId,
update: {
sessionUpdate: "available_commands_update",
availableCommands: await this.#buildAvailableCommands(record.session),
},
});
}
async #emitEndOfTurnUpdates(record: ManagedSessionRecord): Promise<void> {
const sessionId = record.session.sessionId;