Parent directory

index.ts

11545 bytes
  1import type { UserMessage } from "@earendil-works/pi-ai";
  2import type { ExtensionAPI, ExtensionContext } from "@earendil-works/pi-coding-agent";
  3import { createJsonlDecisionAudit } from "../core/audit";
  4import { createJsonlDecisionCache } from "../core/cache";
  5import { parseBash, type BashParseResult } from "../core/bash";
  6import { checkDeterministic, evaluateParsedBash } from "../core/deterministic";
  7import { auditInputSummary, cacheKey, normalizeRequest } from "../core/normalize";
  8import { createPolicyPipeline } from "../core/pipeline";
  9import { createPolicyReviewer, type PiReviewerCaller } from "../core/review";
 10import type { PolicyContext, PolicyEvaluation, PolicyRequest } from "../core/types";
 11import { loadPiPolicyConfig, piPolicyPaths } from "./config";
 12import { PolicyApprovalDialog } from "./policy-approval";
 13import {
 14  createSessionBashAllowOverride,
 15  matchesSessionBashAllowOverride,
 16  restoreSessionBashAllowOverride,
 17  SESSION_BASH_ALLOW_ENTRY,
 18  type SessionBashAllowOverride,
 19} from "./session-override";
 20import { checkStaticPermission } from "./static-permissions";
 21
 22type ToolInput = Record<string, unknown>;
 23
 24function toolInput(input: unknown): ToolInput {
 25  return input && typeof input === "object" && !Array.isArray(input) ? { ...(input as ToolInput) } : {};
 26}
 27
 28function inputSummary(toolName: string, input: ToolInput): string {
 29  if (toolName === "bash" && typeof input.command === "string") return input.command;
 30  return JSON.stringify(input).slice(0, 500);
 31}
 32
 33function approvalMessage(toolName: string, input: ToolInput, result: PolicyEvaluation): string {
 34  const rawResponse =
 35    result.error === "Failed to parse response JSON" && result.rawResponse
 36      ? `\n\nLLM response:\n${result.rawResponse}`
 37      : "";
 38  if (toolName !== "bash" || typeof input.command !== "string") {
 39    return `${toolName}: ${inputSummary(toolName, input)}\n\n${result.decision.reason}${rawResponse}`;
 40  }
 41
 42  const commands = result.approvalRequests
 43    ?.map((request) => request.input.command)
 44    .filter((command): command is string => typeof command === "string");
 45  const requiringReview = commands?.length
 46    ? `\n\nCommands requiring review:\n${commands.map((command) => `- ${command}`).join("\n")}`
 47    : "";
 48  return `Entire Bash input:\n${input.command}${requiringReview}\n\n${result.decision.reason}${rawResponse}`;
 49}
 50
 51function piReviewerCaller(ctx: ExtensionContext): PiReviewerCaller {
 52  return async (
 53    config: {
 54      provider: string;
 55      model: string;
 56      reasoningEffort?: "none" | "minimal" | "low" | "medium" | "high" | "xhigh" | "max";
 57      promptCacheKey?: string;
 58    },
 59    messages: readonly { role: "system" | "user"; content: string }[],
 60    _maxTokens: number,
 61    context: PolicyContext,
 62  ): Promise<string> => {
 63    const model = ctx.modelRegistry.find(config.provider, config.model);
 64    if (!model) throw new Error(`Pi model is unavailable: ${config.provider}/${config.model}`);
 65    const systemPrompt = messages.find((message) => message.role === "system")?.content;
 66    const userMessages: UserMessage[] = messages
 67      .filter((message) => message.role === "user")
 68      .map((message) => ({
 69        role: "user",
 70        content: [{ type: "text", text: message.content }],
 71        timestamp: Date.now(),
 72      }));
 73    const response = await ctx.modelRegistry.complete(
 74      model,
 75      { systemPrompt, messages: userMessages },
 76      {
 77        signal: context.signal,
 78        reasoningEffort: config.reasoningEffort,
 79        sessionId: config.promptCacheKey,
 80      },
 81    );
 82    if (response.stopReason !== "stop")
 83      throw new Error(response.errorMessage ?? `Pi reviewer stopped: ${response.stopReason}`);
 84    const text = response.content
 85      .filter((content): content is { type: "text"; text: string } => content.type === "text")
 86      .map((content) => content.text)
 87      .join("\n");
 88    if (!text.trim()) throw new Error("Pi reviewer returned no text");
 89    return text;
 90  };
 91}
 92
 93function createPiPipeline(
 94  sessionBashAllowOverride: SessionBashAllowOverride | undefined,
 95  cwd: string,
 96  ctx: ExtensionContext,
 97  parsedBash: WeakMap<object, BashParseResult>,
 98  bashReviewCommands: WeakMap<object, string[]>,
 99) {
100  const config = loadPiPolicyConfig(cwd);
101  const paths = piPolicyPaths();
102  return createPolicyPipeline({
103    prechecks: [
104      {
105        source: "session",
106        decide: async (request) => {
107          if (
108            request.toolName !== "bash" ||
109            typeof request.input.command !== "string" ||
110            !matchesSessionBashAllowOverride(sessionBashAllowOverride, request.input.command)
111          )
112            return undefined;
113          return {
114            decision: "allow",
115            reason: "Matched session Bash allow override",
116            category: "session_allow",
117          };
118        },
119      },
120      {
121        source: "static",
122        decide: async (request) => checkStaticPermission(request.toolName, request.input, cwd, config),
123      },
124    ],
125    deterministic: (request, context) => {
126      const parsed = parsedBash.get(request.input);
127      if (!parsed) return checkDeterministic(request.toolName, request.input, context.cwd ?? cwd);
128      const result = evaluateParsedBash(parsed, context.cwd ?? cwd);
129      if (result.reviewCommands) bashReviewCommands.set(request.input, result.reviewCommands);
130      return result.decision;
131    },
132    cache: createJsonlDecisionCache(paths.cacheFile),
133    audit: createJsonlDecisionAudit(paths.auditFile),
134    reviewer: createPolicyReviewer(config.reviewer, piReviewerCaller(ctx), config.externalDirectories),
135    normalize: (request, context) => `${context.cwd ?? ""}\n${normalizeRequest(request.toolName, request.input)}`,
136    cacheKey,
137    inputSummary: auditInputSummary,
138  });
139}
140
141export async function evaluatePiToolCall(
142  toolName: string,
143  input: ToolInput,
144  context: PolicyContext,
145  sessionBashAllowOverride: SessionBashAllowOverride | undefined,
146  parse: (command: string) => Promise<BashParseResult> = parseBash,
147  skipPermissions = false,
148  ctx?: ExtensionContext,
149): Promise<PolicyEvaluation> {
150  if (!ctx) throw new Error("Pi extension context is required");
151  const parsedBash = new WeakMap<object, BashParseResult>();
152  const bashReviewCommands = new WeakMap<object, string[]>();
153  const pipeline = createPiPipeline(
154    sessionBashAllowOverride,
155    context.cwd ?? process.cwd(),
156    ctx,
157    parsedBash,
158    bashReviewCommands,
159  );
160  const evaluationOptions = {
161    skipPrechecks: true,
162    skipLLMReview: skipPermissions,
163  };
164  const request: PolicyRequest = { toolName, input };
165  const prechecked = await pipeline.precheck(request, context);
166  if (prechecked) return prechecked;
167
168  if (toolName !== "bash" || typeof input.command !== "string") {
169    return pipeline.evaluate(request, context, evaluationOptions);
170  }
171
172  const parsed = await parse(input.command);
173  if (parsed.parserUnavailable) {
174    return {
175      decision: {
176        decision: "deny",
177        reason: "Bash parser is unavailable",
178        category: "bash",
179      },
180      source: "deterministic",
181    };
182  }
183
184  parsedBash.set(input, parsed);
185  const result = await pipeline.evaluate(request, context, evaluationOptions);
186  const reviewCommands = bashReviewCommands.get(input);
187  return reviewCommands?.length
188    ? {
189        ...result,
190        approvalRequests: reviewCommands.map((command) => ({ toolName, input: { ...input, command } })),
191      }
192    : result;
193}
194
195export default function policyEngine(pi: ExtensionAPI) {
196  let sessionBashAllowOverride: SessionBashAllowOverride | undefined;
197  let skipPermissions = false;
198
199  pi.on("session_start", (_event, ctx) => {
200    sessionBashAllowOverride = restoreSessionBashAllowOverride(
201      ctx.sessionManager.getBranch(),
202      ctx.sessionManager.getSessionId(),
203    );
204  });
205
206  pi.registerCommand("pe-allow-bash", {
207    description: "Temporarily allow Bash commands that match a regular expression",
208    handler: async (args, ctx) => {
209      const source = args.trim();
210      if (!source) {
211        const status = sessionBashAllowOverride
212          ? `Temporary Bash allow override: /${sessionBashAllowOverride.source}/`
213          : "No temporary Bash allow override. Use /pe-allow-bash <regular expression>.";
214        ctx.ui.notify(status, "info");
215        return;
216      }
217
218      if (source === "clear") {
219        sessionBashAllowOverride = undefined;
220        pi.appendEntry(SESSION_BASH_ALLOW_ENTRY, {
221          pattern: null,
222          sessionId: ctx.sessionManager.getSessionId(),
223        });
224        ctx.ui.notify("Temporary Bash allow override cleared.", "info");
225        return;
226      }
227
228      try {
229        sessionBashAllowOverride = createSessionBashAllowOverride(source);
230      } catch (error) {
231        const message = error instanceof Error ? error.message : "Invalid regular expression";
232        ctx.ui.notify(`Invalid regular expression: ${message}`, "error");
233        return;
234      }
235
236      pi.appendEntry(SESSION_BASH_ALLOW_ENTRY, {
237        pattern: source,
238        sessionId: ctx.sessionManager.getSessionId(),
239      });
240      ctx.ui.notify(`Temporary Bash allow override enabled: /${source}/`, "warning");
241    },
242  });
243
244  pi.registerCommand("pe-skip-permissions", {
245    description: "Temporarily skip LLM permission checks",
246    handler: async (args, ctx) => {
247      if (args.trim() === "clear") {
248        skipPermissions = false;
249        ctx.ui.notify("LLM permission checks enabled.", "info");
250        return;
251      }
252      skipPermissions = true;
253      ctx.ui.notify(
254        "LLM permission checks disabled until this extension reloads. Use /pe-skip-permissions clear to re-enable them.",
255        "warning",
256      );
257    },
258  });
259
260  pi.on("tool_call", async (event, ctx) => {
261    if (event.toolName.startsWith("mcp")) return undefined;
262
263    const input = toolInput(event.input);
264    const result = await evaluatePiToolCall(
265      event.toolName,
266      input,
267      {
268        sessionId: ctx.sessionManager.getSessionId(),
269        cwd: ctx.cwd,
270        signal: ctx.signal,
271      },
272      sessionBashAllowOverride,
273      parseBash,
274      skipPermissions,
275      ctx,
276    );
277
278    if (result.decision.decision === "allow") return undefined;
279    if (result.decision.decision === "deny") {
280      return {
281        block: true,
282        reason: `Denied by policy: ${result.decision.reason}`,
283      };
284    }
285    if (!ctx.hasUI) {
286      return {
287        block: true,
288        reason: `Policy approval required: ${result.decision.reason}`,
289      };
290    }
291
292    pi.events.emit("herdr:blocked", {
293      active: true,
294      label: "Policy approval required",
295    });
296    let approved: boolean;
297    try {
298      const title = "Policy approval required";
299      const message = approvalMessage(event.toolName, input, result);
300      approved =
301        ctx.mode === "tui"
302          ? await ctx.ui.custom(
303              (tui, theme, keybindings, done) =>
304                new PolicyApprovalDialog(tui, theme, keybindings, title, message, done),
305              {
306                overlay: true,
307                overlayOptions: { width: "90%", margin: 1 },
308              },
309            )
310          : await ctx.ui.confirm(title, message);
311    } finally {
312      pi.events.emit("herdr:blocked", { active: false });
313    }
314    if (approved) return undefined;
315    const feedback = await ctx.ui.input(
316      "Why was this rejected?",
317      "Tell the agent why it was rejected and what to do differently",
318    );
319    return {
320      block: true,
321      reason: `Blocked by user: ${feedback?.trim() || result.decision.reason}`,
322    };
323  });
324}