205 lines
9.4 KiB
TypeScript
205 lines
9.4 KiB
TypeScript
import type { ChildProcessWithoutNullStreams } from "node:child_process";
|
|
import { Readable, Writable } from "node:stream";
|
|
import * as acp from "@agentclientprotocol/sdk";
|
|
import type { AgentCapabilities, InitializeResponse, RequestPermissionRequest, RequestPermissionResponse, SessionNotification } from "@agentclientprotocol/sdk";
|
|
import type { PermissionPolicy } from "../config.js";
|
|
|
|
export type SafeActivityCategory = "read" | "write" | "execute" | "search" | "delegate" | "other";
|
|
|
|
export interface AcpClientOptions {
|
|
initializeTimeoutMs: number;
|
|
policy: PermissionPolicy;
|
|
onSessionActivity?: (category?: SafeActivityCategory) => void;
|
|
forbidToolActivity?: boolean;
|
|
onToolViolation?: () => void;
|
|
}
|
|
|
|
export class AcpClient {
|
|
private readonly connection: acp.ClientConnection;
|
|
private capabilities: AgentCapabilities = {};
|
|
private activeSessionId?: string;
|
|
private collecting = false;
|
|
private toolViolation = false;
|
|
private violationReject?: (error: Error) => void;
|
|
private chunks: string[] = [];
|
|
|
|
constructor(private readonly child: ChildProcessWithoutNullStreams, private readonly options: AcpClientOptions) {
|
|
const app = acp.client({ name: "gori-agent" })
|
|
.onRequest(acp.methods.client.session.requestPermission, ({ params }) => {
|
|
if (options.forbidToolActivity) this.recordToolViolation();
|
|
return decidePermission(params, options.policy);
|
|
})
|
|
.onNotification(acp.methods.client.session.update, ({ params }) => this.handleUpdate(params));
|
|
const stream = acp.ndJsonStream(
|
|
Writable.toWeb(child.stdin) as WritableStream<Uint8Array>,
|
|
Readable.toWeb(child.stdout) as ReadableStream<Uint8Array>
|
|
);
|
|
this.connection = app.connect(stream);
|
|
}
|
|
|
|
async initialize(): Promise<InitializeResponse> {
|
|
const response = await withTimeout(this.connection.agent.request(acp.methods.agent.initialize, {
|
|
protocolVersion: acp.PROTOCOL_VERSION,
|
|
clientCapabilities: {},
|
|
clientInfo: { name: "gori-agent", version: "0.1.0" }
|
|
}), this.options.initializeTimeoutMs, "ACP initialize timed out");
|
|
if (response.protocolVersion !== acp.PROTOCOL_VERSION) throw new Error(`Unsupported ACP protocol version ${response.protocolVersion}`);
|
|
this.capabilities = response.agentCapabilities || {};
|
|
return response;
|
|
}
|
|
|
|
async newSession(cwd: string): Promise<string> {
|
|
const response = await withTimeout(
|
|
this.connection.agent.request(acp.methods.agent.session.new, { cwd, mcpServers: [] }),
|
|
this.options.initializeTimeoutMs,
|
|
"ACP session/new timed out"
|
|
);
|
|
this.activeSessionId = response.sessionId;
|
|
if (this.toolViolation) throw new Error("Assistant session attempted forbidden tool activity");
|
|
return response.sessionId;
|
|
}
|
|
|
|
async resumeSession(sessionId: string, cwd: string): Promise<void> {
|
|
this.collecting = false;
|
|
this.chunks = [];
|
|
this.activeSessionId = sessionId;
|
|
const request = <T>(promise: Promise<T>, operation: string): Promise<T> => withTimeout(
|
|
promise,
|
|
this.options.initializeTimeoutMs,
|
|
`ACP session/${operation} timed out`
|
|
);
|
|
try {
|
|
if (this.capabilities.sessionCapabilities?.resume) {
|
|
try {
|
|
await request(this.connection.agent.request(acp.methods.agent.session.resume, { sessionId, cwd, mcpServers: [] }), "resume");
|
|
} catch (error) {
|
|
if (!this.capabilities.loadSession) throw error;
|
|
await request(this.connection.agent.request(acp.methods.agent.session.load, { sessionId, cwd, mcpServers: [] }), "load");
|
|
}
|
|
} else if (this.capabilities.loadSession) {
|
|
await request(this.connection.agent.request(acp.methods.agent.session.load, { sessionId, cwd, mcpServers: [] }), "load");
|
|
} else {
|
|
throw new Error("ACP backend cannot resume or load sessions");
|
|
}
|
|
} catch (error) {
|
|
this.activeSessionId = undefined;
|
|
throw error;
|
|
}
|
|
this.chunks = [];
|
|
}
|
|
|
|
async prompt(text: string, cancellationSignal?: AbortSignal): Promise<string> {
|
|
if (!this.activeSessionId) throw new Error("ACP session is not active");
|
|
if (this.toolViolation) throw new Error("Assistant session attempted forbidden tool activity");
|
|
this.chunks = [];
|
|
this.collecting = true;
|
|
const violation = new Promise<never>((_resolve, reject) => { this.violationReject = reject; });
|
|
try {
|
|
await Promise.race([
|
|
this.connection.agent.request(acp.methods.agent.session.prompt, {
|
|
sessionId: this.activeSessionId,
|
|
prompt: [{ type: "text", text }]
|
|
}, cancellationSignal ? { cancellationSignal } : undefined),
|
|
violation
|
|
]);
|
|
if (this.toolViolation) throw new Error("Assistant session attempted forbidden tool activity");
|
|
return this.chunks.join("").trim();
|
|
} finally {
|
|
this.violationReject = undefined;
|
|
this.collecting = false;
|
|
}
|
|
}
|
|
|
|
async settleIsolation(): Promise<void> {
|
|
if (!this.options.forbidToolActivity) return;
|
|
await new Promise((resolve) => setTimeout(resolve, 50));
|
|
if (this.toolViolation) throw new Error("Assistant session attempted forbidden tool activity");
|
|
}
|
|
|
|
async cancel(): Promise<void> {
|
|
if (this.activeSessionId) await this.connection.agent.notify(acp.methods.agent.session.cancel, { sessionId: this.activeSessionId });
|
|
}
|
|
|
|
async closeSession(): Promise<void> {
|
|
if (this.activeSessionId && this.capabilities.sessionCapabilities?.close) {
|
|
await this.connection.agent.request(acp.methods.agent.session.close, { sessionId: this.activeSessionId }).catch(() => undefined);
|
|
}
|
|
}
|
|
|
|
close(error?: unknown): void { this.connection.close(error); }
|
|
|
|
private handleUpdate(notification: SessionNotification): void {
|
|
if (this.activeSessionId && notification.sessionId !== this.activeSessionId) return;
|
|
const update = notification.update;
|
|
const category = activityCategory(update);
|
|
if (category && this.options.forbidToolActivity) this.recordToolViolation();
|
|
this.options.onSessionActivity?.(category);
|
|
if (!this.collecting || notification.sessionId !== this.activeSessionId) return;
|
|
if (update.sessionUpdate === "agent_message_chunk" && update.content.type === "text") this.chunks.push(update.content.text);
|
|
}
|
|
|
|
private recordToolViolation(): void {
|
|
if (this.toolViolation) return;
|
|
this.toolViolation = true;
|
|
this.violationReject?.(new Error("Assistant session attempted forbidden tool activity"));
|
|
this.options.onToolViolation?.();
|
|
}
|
|
}
|
|
|
|
function activityCategory(update: SessionNotification["update"]): SafeActivityCategory | undefined {
|
|
if (!String(update.sessionUpdate).startsWith("tool_call")) return undefined;
|
|
const record = update as unknown as Record<string, unknown>;
|
|
const raw = JSON.stringify({ kind: record.kind, name: record.name, title: record.title }).toLowerCase();
|
|
if (/read|view|fetch/.test(raw)) return "read";
|
|
if (/grep|glob|search|find/.test(raw)) return "search";
|
|
if (/write|edit|patch/.test(raw)) return "write";
|
|
if (/bash|terminal|execute|command/.test(raw)) return "execute";
|
|
if (/agent|delegate|subagent/.test(raw)) return "delegate";
|
|
return "other";
|
|
}
|
|
|
|
export function decidePermission(request: RequestPermissionRequest, policy: PermissionPolicy): RequestPermissionResponse {
|
|
if (policy.mode === "deny") return reject(request);
|
|
const allowOption = request.options.find((option) => option.kind === "allow_once") || request.options.find((option) => option.kind === "allow_always");
|
|
if (!allowOption) return reject(request);
|
|
if (policy.mode === "auto") return { outcome: { outcome: "selected", optionId: allowOption.optionId } };
|
|
|
|
const name = typeof request.toolCall.name === "string" ? request.toolCall.name.trim().toLowerCase() : "";
|
|
if (!name || !policy.allowedTools.some((tool) => name === tool.trim().toLowerCase())) return reject(request);
|
|
if (name === "bash" || name === "terminal") {
|
|
if (policy.allowedCommandPatterns.length === 0) return reject(request);
|
|
const input = commandInput(request.toolCall.rawInput);
|
|
if (!input || !policy.allowedCommandPatterns.some((pattern) => fullMatch(pattern, input))) return reject(request);
|
|
}
|
|
return { outcome: { outcome: "selected", optionId: allowOption.optionId } };
|
|
}
|
|
|
|
function commandInput(rawInput: unknown): string | undefined {
|
|
if (typeof rawInput === "string") return rawInput;
|
|
if (typeof rawInput === "object" && rawInput !== null && !Array.isArray(rawInput)) {
|
|
const record = rawInput as Record<string, unknown>;
|
|
if (Object.keys(record).some((key) => !["command", "timeout", "timeoutMs"].includes(key))) return undefined;
|
|
return typeof record.command === "string" ? record.command : undefined;
|
|
}
|
|
return undefined;
|
|
}
|
|
|
|
function fullMatch(pattern: string, input: string): boolean {
|
|
const match = new RegExp(pattern).exec(input);
|
|
return match?.index === 0 && match[0] === input;
|
|
}
|
|
|
|
function reject(request: RequestPermissionRequest): RequestPermissionResponse {
|
|
const option = request.options.find((item) => item.kind === "reject_once") || request.options.find((item) => item.kind === "reject_always");
|
|
return option ? { outcome: { outcome: "selected", optionId: option.optionId } } : { outcome: { outcome: "cancelled" } };
|
|
}
|
|
|
|
async function withTimeout<T>(promise: Promise<T>, timeoutMs: number, message: string): Promise<T> {
|
|
let timer: NodeJS.Timeout | undefined;
|
|
try {
|
|
return await Promise.race([promise, new Promise<T>((_resolve, reject) => { timer = setTimeout(() => reject(new Error(message)), timeoutMs); })]);
|
|
} finally {
|
|
if (timer) clearTimeout(timer);
|
|
}
|
|
}
|