fix: added Cursor OAuth provider and fixed authentication flow race conditions

- Added Cursor OAuth provider support to CLI authentication flow.
- Fixed race conditions in OAuth callback handling and session picker cleanup.
- Improved JWT token parsing and error handling across OAuth providers.
- Replaced Bun.sleep() with abortable sleep for proper cancellation support.
- Enhanced process execution with concurrent stream reading and better timeout handling.
- Added cache cleanup mechanism for auth storage with 5-minute intervals.
This commit is contained in:
can1357
2026-01-20 17:07:02 +01:00
parent 94f406d4f0
commit 21af96d890
35 changed files with 468 additions and 287 deletions
+6 -7
View File
@@ -105,20 +105,19 @@ export function streamProxy(model: Model<any>, context: Context, options: ProxyS
timestamp: Date.now(),
};
let reader: ReadableStreamDefaultReader<Uint8Array> | undefined;
let response: Response | null = null;
const abortHandler = () => {
if (reader) {
reader.cancel("Request aborted by user").catch(() => {});
const body = response?.body;
if (body) {
body.cancel("Request aborted by user").catch(() => {});
}
};
if (options.signal) {
options.signal.addEventListener("abort", abortHandler);
options.signal.addEventListener("abort", abortHandler, { once: true });
}
try {
const response = await fetch(`${options.proxyUrl}/api/stream`, {
response = await fetch(`${options.proxyUrl}/api/stream`, {
method: "POST",
headers: {
Authorization: `Bearer ${options.authToken}`,
+12
View File
@@ -3,6 +3,7 @@ import { createInterface } from "readline";
import { CliAuthStorage } from "./storage";
import "./utils/migrate-env";
import { loginAnthropic } from "./utils/oauth/anthropic";
import { loginCursor } from "./utils/oauth/cursor";
import { loginGitHubCopilot } from "./utils/oauth/github-copilot";
import { loginAntigravity } from "./utils/oauth/google-antigravity";
import { loginGeminiCli } from "./utils/oauth/google-gemini-cli";
@@ -88,6 +89,17 @@ async function login(provider: OAuthProvider): Promise<void> {
});
break;
case "cursor":
credentials = await loginCursor(
(url) => {
console.log(`\nOpen this URL in your browser:\n${url}\n`);
},
() => {
console.log("Waiting for browser authentication...");
},
);
break;
default:
throw new Error(`Unknown provider: ${provider}`);
}
@@ -1,4 +1,5 @@
import os from "node:os";
import { abortableSleep } from "@oh-my-pi/pi-utils";
import type {
ResponseFunctionToolCall,
ResponseInput,
@@ -440,13 +441,13 @@ async function fetchWithRetry(url: string, init: RequestInit, signal?: AbortSign
}
if (signal?.aborted) return response;
const delay = getRetryDelayMs(response, attempt);
await Bun.sleep(delay);
await abortableSleep(delay, signal);
} catch (error) {
if (attempt >= CODEX_MAX_RETRIES || signal?.aborted) {
throw error;
}
const delay = CODEX_RETRY_DELAY_MS * (attempt + 1);
await Bun.sleep(delay);
await abortableSleep(delay, signal);
}
attempt += 1;
}
+6 -1
View File
@@ -63,7 +63,12 @@ class AnthropicOAuthFlow extends OAuthCallbackFlow {
});
if (!tokenResponse.ok) {
const error = await tokenResponse.text();
let error: string;
try {
error = await tokenResponse.text();
} catch {
error = `HTTP ${tokenResponse.status}`;
}
throw new Error(`Token exchange failed: ${error}`);
}
@@ -158,12 +158,14 @@ export abstract class OAuthCallbackFlow {
resultState = { ok: true, code, state };
}
// Signal to waitForCallback
// Signal to waitForCallback - capture refs before they could be cleared
const resolve = this.callbackResolve;
const reject = this.callbackReject;
queueMicrotask(() => {
if (resultState.ok) {
this.callbackResolve?.({ code: resultState.code, state: resultState.state });
resolve?.({ code: resultState.code, state: resultState.state });
} else {
this.callbackReject?.(resultState.error ?? "Unknown error");
reject?.(resultState.error ?? "Unknown error");
}
});
+6 -2
View File
@@ -126,12 +126,16 @@ export async function refreshCursorToken(apiKeyOrRefreshToken: string): Promise<
function getTokenExpiry(token: string): number {
try {
const [, payload] = token.split(".");
const parts = token.split(".");
if (parts.length !== 3) {
return Date.now() + 3600 * 1000;
}
const payload = parts[1];
if (!payload) {
return Date.now() + 3600 * 1000;
}
const decoded = JSON.parse(atob(payload.replace(/-/g, "+").replace(/_/g, "/")));
if (decoded.exp) {
if (decoded && typeof decoded === "object" && typeof decoded.exp === "number") {
return decoded.exp * 1000 - 5 * 60 * 1000;
}
} catch {
+55 -89
View File
@@ -40,11 +40,16 @@ function getAccountId(accessToken: string): string | null {
return typeof accountId === "string" && accountId.length > 0 ? accountId : null;
}
class OpenAICodexOAuthFlow extends OAuthCallbackFlow {
private verifier: string = "";
private challenge: string = "";
interface PKCE {
verifier: string;
challenge: string;
}
constructor(ctrl: OAuthController) {
class OpenAICodexOAuthFlow extends OAuthCallbackFlow {
constructor(
ctrl: OAuthController,
private readonly pkce: PKCE,
) {
super(ctrl, CALLBACK_PORT, CALLBACK_PATH);
}
@@ -52,16 +57,12 @@ class OpenAICodexOAuthFlow extends OAuthCallbackFlow {
state: string,
redirectUri: string,
): Promise<{ url: string; instructions?: string }> {
const pkce = await generatePKCE();
this.verifier = pkce.verifier;
this.challenge = pkce.challenge;
const searchParams = new URLSearchParams({
response_type: "code",
client_id: CLIENT_ID,
redirect_uri: redirectUri,
scope: SCOPE,
code_challenge: this.challenge,
code_challenge: this.pkce.challenge,
code_challenge_method: "S256",
state,
id_token_add_organizations: "true",
@@ -74,56 +75,61 @@ class OpenAICodexOAuthFlow extends OAuthCallbackFlow {
}
protected async exchangeToken(code: string, _state: string, redirectUri: string): Promise<OAuthCredentials> {
const tokenResponse = await fetch(TOKEN_URL, {
method: "POST",
headers: { "Content-Type": "application/x-www-form-urlencoded" },
body: new URLSearchParams({
grant_type: "authorization_code",
client_id: CLIENT_ID,
code,
code_verifier: this.verifier,
redirect_uri: redirectUri,
}),
});
if (!tokenResponse.ok) {
throw new Error(`Token exchange failed: ${tokenResponse.status}`);
}
const tokenData = (await tokenResponse.json()) as {
access_token?: string;
refresh_token?: string;
expires_in?: number;
};
if (!tokenData.access_token || !tokenData.refresh_token || typeof tokenData.expires_in !== "number") {
throw new Error("Token response missing required fields");
}
const accountId = getAccountId(tokenData.access_token);
if (!accountId) {
throw new Error("Failed to extract accountId from token");
}
return {
access: tokenData.access_token,
refresh: tokenData.refresh_token,
expires: Date.now() + tokenData.expires_in * 1000,
accountId,
};
return exchangeCodeForToken(code, this.pkce.verifier, redirectUri);
}
}
async function exchangeCodeForToken(code: string, verifier: string, redirectUri: string): Promise<OAuthCredentials> {
const tokenResponse = await fetch(TOKEN_URL, {
method: "POST",
headers: { "Content-Type": "application/x-www-form-urlencoded" },
body: new URLSearchParams({
grant_type: "authorization_code",
client_id: CLIENT_ID,
code,
code_verifier: verifier,
redirect_uri: redirectUri,
}),
});
if (!tokenResponse.ok) {
throw new Error(`Token exchange failed: ${tokenResponse.status}`);
}
const tokenData = (await tokenResponse.json()) as {
access_token?: string;
refresh_token?: string;
expires_in?: number;
};
if (!tokenData.access_token || !tokenData.refresh_token || typeof tokenData.expires_in !== "number") {
throw new Error("Token response missing required fields");
}
const accountId = getAccountId(tokenData.access_token);
if (!accountId) {
throw new Error("Failed to extract accountId from token");
}
return {
access: tokenData.access_token,
refresh: tokenData.refresh_token,
expires: Date.now() + tokenData.expires_in * 1000,
accountId,
};
}
/**
* Login with OpenAI Codex OAuth
*/
export async function loginOpenAICodex(ctrl: OAuthController): Promise<OAuthCredentials> {
const flow = new OpenAICodexOAuthFlow(ctrl);
const pkce = await generatePKCE();
const flow = new OpenAICodexOAuthFlow(ctrl, pkce);
const redirectUri = `http://localhost:${CALLBACK_PORT}${CALLBACK_PATH}`;
try {
return await flow.login();
} catch (error) {
// Callback failed - fall back to onPrompt if available
if (!ctrl.onPrompt) {
throw error;
}
@@ -139,47 +145,7 @@ export async function loginOpenAICodex(ctrl: OAuthController): Promise<OAuthCred
throw new Error("No authorization code found in input");
}
const redirectUri = `http://localhost:${CALLBACK_PORT}${CALLBACK_PATH}`;
// Manual token exchange
const pkce = await generatePKCE();
const tokenResponse = await fetch(TOKEN_URL, {
method: "POST",
headers: { "Content-Type": "application/x-www-form-urlencoded" },
body: new URLSearchParams({
grant_type: "authorization_code",
client_id: CLIENT_ID,
code: parsed.code,
code_verifier: pkce.verifier,
redirect_uri: redirectUri,
}),
});
if (!tokenResponse.ok) {
throw new Error(`Token exchange failed: ${tokenResponse.status}`);
}
const tokenData = (await tokenResponse.json()) as {
access_token?: string;
refresh_token?: string;
expires_in?: number;
};
if (!tokenData.access_token || !tokenData.refresh_token || typeof tokenData.expires_in !== "number") {
throw new Error("Token response missing required fields");
}
const accountId = getAccountId(tokenData.access_token);
if (!accountId) {
throw new Error("Failed to extract accountId from token");
}
return {
access: tokenData.access_token,
refresh: tokenData.refresh_token,
expires: Date.now() + tokenData.expires_in * 1000,
accountId,
};
return exchangeCodeForToken(parsed.code, pkce.verifier, redirectUri);
}
}
+4
View File
@@ -2,6 +2,10 @@
## [Unreleased]
### Fixed
- Fixed unhandled promise rejection when tool execution fails by adding missing `.catch()` to floating `.finally()` chain in `createAbortablePromise`
## [6.8.0] - 2026-01-20
### Added
@@ -247,16 +247,16 @@ export default function planModeExtension(pi: ExtensionAPI) {
}
}
function togglePlanMode(ctx: ExtensionContext) {
async function togglePlanMode(ctx: ExtensionContext) {
planModeEnabled = !planModeEnabled;
executionMode = false;
todoItems = [];
if (planModeEnabled) {
pi.setActiveTools(PLAN_MODE_TOOLS);
await pi.setActiveTools(PLAN_MODE_TOOLS);
ctx.ui.notify(`Plan mode enabled. Tools: ${PLAN_MODE_TOOLS.join(", ")}`);
} else {
pi.setActiveTools(NORMAL_MODE_TOOLS);
await pi.setActiveTools(NORMAL_MODE_TOOLS);
ctx.ui.notify("Plan mode disabled. Full access restored.");
}
updateStatus(ctx);
@@ -266,7 +266,7 @@ export default function planModeExtension(pi: ExtensionAPI) {
pi.registerCommand("plan", {
description: "Toggle plan mode (read-only exploration)",
handler: async (_args, ctx) => {
togglePlanMode(ctx);
await togglePlanMode(ctx);
},
});
@@ -294,7 +294,7 @@ export default function planModeExtension(pi: ExtensionAPI) {
pi.registerShortcut(Key.shift("p"), {
description: "Toggle plan mode",
handler: async (ctx) => {
togglePlanMode(ctx);
await togglePlanMode(ctx);
},
});
@@ -417,7 +417,7 @@ Execute each step in order.`,
executionMode = false;
todoItems = [];
pi.setActiveTools(NORMAL_MODE_TOOLS);
await pi.setActiveTools(NORMAL_MODE_TOOLS);
updateStatus(ctx);
}
return;
@@ -470,7 +470,7 @@ Execute each step in order.`,
if (choice?.startsWith("Execute")) {
planModeEnabled = false;
executionMode = hasTodos;
pi.setActiveTools(NORMAL_MODE_TOOLS);
await pi.setActiveTools(NORMAL_MODE_TOOLS);
updateStatus(ctx);
// Simple execution message - context event filters old plan mode messages
@@ -519,7 +519,7 @@ Execute each step in order.`,
}
if (planModeEnabled) {
pi.setActiveTools(PLAN_MODE_TOOLS);
await pi.setActiveTools(PLAN_MODE_TOOLS);
}
updateStatus(ctx);
});
@@ -31,12 +31,12 @@ export default function toolsExtension(pi: ExtensionAPI) {
}
// Apply current tool selection
function applyTools() {
pi.setActiveTools(Array.from(enabledTools));
async function applyTools() {
await pi.setActiveTools(Array.from(enabledTools));
}
// Find the last tools-config entry in the current branch
function restoreFromBranch(ctx: ExtensionContext) {
async function restoreFromBranch(ctx: ExtensionContext) {
allTools = pi.getAllTools();
// Get entries in current branch only
@@ -55,7 +55,7 @@ export default function toolsExtension(pi: ExtensionAPI) {
if (savedTools) {
// Restore saved tool selection (filter to only tools that still exist)
enabledTools = new Set(savedTools.filter((t: string) => allTools.includes(t)));
applyTools();
await applyTools();
} else {
// No saved state - sync with currently active tools
enabledTools = new Set(pi.getActiveTools());
@@ -130,16 +130,16 @@ export default function toolsExtension(pi: ExtensionAPI) {
// Restore state on session start
pi.on("session_start", async (_event, ctx) => {
restoreFromBranch(ctx);
await restoreFromBranch(ctx);
});
// Restore state when navigating the session tree
pi.on("session_tree", async (_event, ctx) => {
restoreFromBranch(ctx);
await restoreFromBranch(ctx);
});
// Restore state after branching
pi.on("session_branch", async (_event, ctx) => {
restoreFromBranch(ctx);
await restoreFromBranch(ctx);
});
}
@@ -28,8 +28,11 @@ export async function selectSession(sessions: SessionInfo[]): Promise<string | n
}
},
() => {
ui.stop();
process.exit(0);
if (!resolved) {
resolved = true;
ui.stop();
process.exit(0);
}
},
);
@@ -2226,7 +2226,7 @@ export class AgentSession {
error: message,
model: `${candidate.provider}/${candidate.id}`,
});
await Bun.sleep(delayMs);
await abortableSleep(delayMs, this._autoCompactionAbortController.signal);
}
}
@@ -2291,6 +2291,10 @@ export class AgentSession {
}, 100);
}
} catch (error) {
if (this._autoCompactionAbortController?.signal.aborted) {
this._emit({ type: "auto_compaction_end", result: undefined, aborted: true, willRetry: false });
return;
}
const errorMessage = error instanceof Error ? error.message : "compaction failed";
this._emit({
type: "auto_compaction_end",
@@ -107,10 +107,12 @@ function toBoolean(value: unknown): boolean | undefined {
export class AuthStorage {
private static readonly codexUsageCacheTtlMs = 60_000; // Cache usage data for 1 minute
private static readonly defaultBackoffMs = 60_000; // Default backoff when no reset time available
private static readonly cacheCleanupIntervalMs = 300_000; // Clean expired cache every 5 minutes
/** Provider -> credentials cache, populated from agent.db on reload(). */
private data: Map<string, StoredCredential[]> = new Map();
private storage: AgentStorage;
private lastCacheCleanup = 0;
/** Resolved path to agent.db (derived from authPath or used directly if .db). */
private dbPath: string;
private runtimeOverrides: Map<string, string> = new Map();
@@ -153,6 +155,7 @@ export class AuthStorage {
instance.sessionLastCredential = new Map();
instance.credentialBackoff = new Map();
instance.codexUsageCache = new Map();
instance.lastCacheCleanup = 0;
for (const [provider, creds] of Object.entries(data.credentials)) {
instance.data.set(
@@ -748,6 +751,11 @@ export class AuthStorage {
const cacheKey = this.getCodexUsageCacheKey(accountId, normalizedBase);
const now = Date.now();
if (now - this.lastCacheCleanup > AuthStorage.cacheCleanupIntervalMs) {
this.lastCacheCleanup = now;
this.storage.cleanExpiredCache();
}
// Check in-memory cache first (fastest)
const memCached = this.codexUsageCache.get(cacheKey);
if (memCached && memCached.expiresAt > now) {
@@ -4,7 +4,7 @@
* Provides unified bash execution for AgentSession.executeBash() and direct calls.
*/
import { cspawn, Exception } from "@oh-my-pi/pi-utils";
import { cspawn, Exception, ptree } from "@oh-my-pi/pi-utils";
import { getShellConfig } from "../utils/shell";
import { getOrCreateSnapshot, getSnapshotSourceCommand } from "../utils/shell-snapshot";
import { OutputSink } from "./streaming-output";
@@ -63,7 +63,7 @@ export async function executeBash(command: string, options?: BashExecutorOptions
// Exception covers NonZeroExitError, AbortError, TimeoutError
if (err instanceof Exception) {
if (err.aborted) {
const isTimeout = err.message.includes("timed out");
const isTimeout = err instanceof ptree.TimeoutError || err.message.toLowerCase().includes("timed out");
const annotation = isTimeout
? `Command timed out after ${Math.round((options?.timeout ?? 0) / 1000)} seconds`
: undefined;
+4 -2
View File
@@ -41,14 +41,16 @@ export async function execCommand(
signal: options?.signal,
timeout: options?.timeout,
});
// Read streams before awaiting exit to avoid data loss if streams close
const [stdoutText, stderrText] = await Promise.all([proc.stdout.text(), proc.stderr.text()]);
try {
await proc.exited;
} catch {
// ChildProcess rejects on non-zero exit; we handle it below
}
return {
stdout: await proc.stdout.text(),
stderr: await proc.stderr.text(),
stdout: stdoutText,
stderr: stderrText,
code: proc.exitCode ?? 0,
killed: proc.exitReason instanceof ptree.AbortError,
};
@@ -757,7 +757,7 @@ export interface ExtensionAPI {
getAllTools(): string[];
/** Set the active tools by name. */
setActiveTools(toolNames: string[]): void;
setActiveTools(toolNames: string[]): Promise<void>;
/** Set the current model. Returns false if no API key available. */
setModel(model: Model<any>): Promise<boolean>;
@@ -685,12 +685,13 @@ export interface HookAPI {
* @param message.content - Message content (string or TextContent/ImageContent array)
* @param message.display - Whether to show in TUI (true = styled display, false = hidden)
* @param message.details - Optional hook-specific metadata (not sent to LLM)
* @param triggerTurn - If true and agent is idle, triggers a new LLM turn. Default: false.
* If agent is streaming, message is queued and triggerTurn is ignored.
* @param options.triggerTurn - If true and agent is idle, triggers a new LLM turn. Default: false.
* If agent is streaming, message is queued and triggerTurn is ignored.
* @param options.deliverAs - How to deliver the message: "steer" or "followUp".
*/
sendMessage<T = unknown>(
message: Pick<HookMessage<T>, "customType" | "content" | "display" | "details">,
triggerTurn?: boolean,
options?: { triggerTurn?: boolean; deliverAs?: "steer" | "followUp" },
): void;
/**
@@ -5,7 +5,7 @@
* Based on MCP spec 2025-03-26.
*/
import type { JsonRpcResponse, MCPHttpServerConfig, MCPSseServerConfig, MCPTransport } from "../types";
import type { JsonRpcMessage, JsonRpcResponse, MCPHttpServerConfig, MCPSseServerConfig, MCPTransport } from "../types";
/** Generate unique request ID */
function generateId(): string {
@@ -167,39 +167,47 @@ export class HttpTransport implements MCPTransport {
throw new Error("No response body");
}
let result: T | undefined;
const timeout = this.config.timeout ?? 30000;
for await (const event of readSseEvents(response.body)) {
const data = event.data?.trim();
if (!data || data === "[DONE]") continue;
try {
const message = JSON.parse(data) as JsonRpcResponse;
const parse = async (): Promise<T> => {
for await (const event of readSseEvents(response.body!)) {
const data = event.data?.trim();
if (!data || data === "[DONE]") continue;
// Handle our response
if ("id" in message && message.id === expectedId) {
if (message.error) {
throw new Error(`MCP error ${message.error.code}: ${message.error.message}`);
try {
const message = JSON.parse(data) as JsonRpcMessage;
if (
"id" in message &&
(message as JsonRpcResponse).id === expectedId &&
("result" in message || "error" in message)
) {
const response = message as JsonRpcResponse;
if (response.error) {
throw new Error(`MCP error ${response.error.code}: ${response.error.message}`);
}
return response.result as T;
}
if ("method" in message && !("id" in message)) {
this.onNotification?.(message.method, message.params);
}
} catch (error) {
if (error instanceof Error && error.message.startsWith("MCP error")) {
throw error;
}
result = message.result as T;
}
// Handle notifications
else if ("method" in message && !("id" in message)) {
const notification = message as { method: string; params?: unknown };
this.onNotification?.(notification.method, notification.params);
}
} catch (error) {
if (error instanceof Error && error.message.startsWith("MCP error")) {
throw error;
}
// Ignore other parse errors
}
}
if (result === undefined) {
throw new Error("No response received");
}
throw new Error(`No response received for request ID ${expectedId}`);
};
return result;
return Promise.race([
parse(),
new Promise<never>((_, reject) =>
setTimeout(() => reject(new Error(`SSE response timeout after ${timeout}ms`)), timeout),
),
]);
}
async notify(method: string, params?: Record<string, unknown>): Promise<void> {
@@ -116,8 +116,8 @@ const DEFAULT_ENV_DENYLIST = new Set([
const CASE_INSENSITIVE_ENV = process.platform === "win32";
const ACTIVE_ENV_ALLOWLIST = CASE_INSENSITIVE_ENV ? WINDOWS_ENV_ALLOWLIST : DEFAULT_ENV_ALLOWLIST;
const NORMALIZED_ALLOWLIST = new Set(
Array.from(ACTIVE_ENV_ALLOWLIST, (key) => (CASE_INSENSITIVE_ENV ? key.toUpperCase() : key)),
const NORMALIZED_ALLOWLIST = new Map(
Array.from(ACTIVE_ENV_ALLOWLIST, (key) => [CASE_INSENSITIVE_ENV ? key.toUpperCase() : key, key] as const),
);
const NORMALIZED_DENYLIST = new Set(
Array.from(DEFAULT_ENV_DENYLIST, (key) => (CASE_INSENSITIVE_ENV ? key.toUpperCase() : key)),
@@ -168,8 +168,9 @@ function filterEnv(env: Record<string, string | undefined>): Record<string, stri
if (value === undefined) continue;
const normalizedKey = normalizeEnvKey(key);
if (NORMALIZED_DENYLIST.has(normalizedKey)) continue;
if (NORMALIZED_ALLOWLIST.has(normalizedKey)) {
filtered[key] = value;
const canonicalKey = NORMALIZED_ALLOWLIST.get(normalizedKey);
if (canonicalKey !== undefined) {
filtered[canonicalKey] = value;
continue;
}
if (NORMALIZED_ALLOW_PREFIXES.some((prefix) => normalizedKey.startsWith(prefix))) {
@@ -95,7 +95,7 @@ export async function executeSSH(
return {
exitCode: undefined,
cancelled: true,
...sink.dump(`SSH command timed out after ${Math.round(options!.timeout! / 1000)} seconds`),
...sink.dump(`SSH: ${err.message}`),
};
}
if (err.aborted) {
@@ -64,6 +64,7 @@ export interface PythonToolDetails {
images?: ImageContent[];
/** Structured status events from prelude helpers */
statusEvents?: PythonStatusEvent[];
isError?: boolean;
}
function formatJsonScalar(value: unknown): string {
@@ -343,7 +343,9 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
let abortSent = false;
let abortReason: AbortReason | undefined;
let terminationScheduled = false;
let pendingTerminationController: AbortController | null = null;
let terminated = false;
let terminationTimeoutId: ReturnType<typeof setTimeout> | null = null;
let pendingTerminationTimeoutId: ReturnType<typeof setTimeout> | null = null;
let finalize: ((message: Extract<SubagentWorkerResponse, { type: "done" }>) => void) | null = null;
const listenerController = new AbortController();
const listenerSignal = listenerController.signal;
@@ -416,28 +418,25 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
const scheduleTermination = () => {
if (terminationScheduled) return;
terminationScheduled = true;
const timeoutSignal = AbortSignal.timeout(2000);
timeoutSignal.addEventListener(
"abort",
() => {
if (resolved) return;
try {
worker.terminate();
} catch {
// Ignore termination errors
}
if (finalize && !resolved) {
finalize({
type: "done",
exitCode: 1,
durationMs: Date.now() - startTime,
error: abortReason === "signal" ? "Aborted" : "Worker terminated after tool completion",
aborted: abortReason === "signal",
});
}
},
{ once: true, signal: listenerSignal },
);
terminationTimeoutId = setTimeout(() => {
terminationTimeoutId = null;
if (resolved || terminated) return;
terminated = true;
try {
worker.terminate();
} catch {
// Ignore termination errors
}
if (finalize && !resolved) {
finalize({
type: "done",
exitCode: 1,
durationMs: Date.now() - startTime,
error: abortReason === "signal" ? "Aborted" : "Worker terminated after tool completion",
aborted: abortReason === "signal",
});
}
}, 2000);
};
const requestAbort = (reason: AbortReason) => {
@@ -461,28 +460,25 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
// Worker already terminated, nothing to do
}
// Cancel pending termination if it exists
if (pendingTerminationController) {
pendingTerminationController.abort();
pendingTerminationController = null;
}
cancelPendingTermination();
scheduleTermination();
};
const schedulePendingTermination = () => {
if (pendingTerminationController || abortSent || terminationScheduled || resolved) return;
const readyController = new AbortController();
pendingTerminationController = readyController;
const pendingSignal = AbortSignal.any([AbortSignal.timeout(2000), readyController.signal]);
pendingSignal.addEventListener(
"abort",
() => {
pendingTerminationController = null;
if (!resolved) {
requestAbort("terminate");
}
},
{ once: true, signal: listenerSignal },
);
if (pendingTerminationTimeoutId || abortSent || terminationScheduled || resolved) return;
pendingTerminationTimeoutId = setTimeout(() => {
pendingTerminationTimeoutId = null;
if (!resolved) {
requestAbort("terminate");
}
}, 2000);
};
const cancelPendingTermination = () => {
if (pendingTerminationTimeoutId) {
clearTimeout(pendingTerminationTimeoutId);
pendingTerminationTimeoutId = null;
}
};
// Handle abort signal
@@ -655,9 +651,10 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
// Accumulate tokens for progress display
progress.tokens += getUsageTokens(messageUsage);
}
// If pending termination, now we have tokens - terminate
if (pendingTerminationController) {
pendingTerminationController.abort();
// If pending termination, now we have tokens - terminate immediately
if (pendingTerminationTimeoutId) {
cancelPendingTermination();
requestAbort("terminate");
}
break;
}
@@ -714,7 +711,6 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
const done = await new Promise<Extract<SubagentWorkerResponse, { type: "done" }>>((resolve) => {
const cleanup = () => {
pendingTerminationController = null;
listenerController.abort();
};
finalize = (message) => {
@@ -723,10 +719,18 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
cleanup();
resolve(message);
};
const postMessageSafe = (message: unknown) => {
if (resolved || terminated) return;
try {
worker.postMessage(message);
} catch {
// Worker already terminated
}
};
const handleMCPCall = async (request: MCPToolCallRequest) => {
const mcpManager = options.mcpManager;
if (!mcpManager) {
worker.postMessage({
postMessageSafe({
type: "mcp_tool_result",
callId: request.callId,
error: "MCP not available",
@@ -743,13 +747,13 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
})(),
request.timeoutMs,
);
worker.postMessage({
postMessageSafe({
type: "mcp_tool_result",
callId: request.callId,
result: { content: result.content ?? [], isError: result.isError },
});
} catch (error) {
worker.postMessage({
postMessageSafe({
type: "mcp_tool_result",
callId: request.callId,
error: error instanceof Error ? error.message : String(error),
@@ -767,7 +771,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
const handlePythonCall = async (request: PythonToolCallRequest) => {
if (!pythonTool) {
worker.postMessage({
postMessageSafe({
type: "python_tool_result",
callId: request.callId,
error: "Python proxy not available",
@@ -785,7 +789,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
request.params as { code: string; timeout?: number; workdir?: string; reset?: boolean },
combinedSignal,
);
worker.postMessage({
postMessageSafe({
type: "python_tool_result",
callId: request.callId,
result: { content: result.content ?? [], details: result.details },
@@ -797,7 +801,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
: error instanceof Error
? error.message
: String(error);
worker.postMessage({
postMessageSafe({
type: "python_tool_result",
callId: request.callId,
error: message,
@@ -816,7 +820,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
const handleLspCall = async (request: LspToolCallRequest) => {
if (!lspTool) {
worker.postMessage({
postMessageSafe({
type: "lsp_tool_result",
callId: request.callId,
error: "LSP proxy not available",
@@ -828,7 +832,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
lspTool.execute(request.callId, request.params as LspParams, signal),
request.timeoutMs,
);
worker.postMessage({
postMessageSafe({
type: "lsp_tool_result",
callId: request.callId,
result: { content: result.content ?? [], details: result.details },
@@ -840,7 +844,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
: error instanceof Error
? error.message
: String(error);
worker.postMessage({
postMessageSafe({
type: "lsp_tool_result",
callId: request.callId,
error: message,
@@ -881,10 +885,14 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
return;
}
if (message.type === "done") {
// Worker is exiting - mark as terminated to prevent calling terminate() on dead worker
terminated = true;
finalize?.(message);
}
};
const onError = (event: WorkerErrorEvent) => {
// Worker error likely means it's dead or dying
terminated = true;
finalize?.({
type: "done",
exitCode: 1,
@@ -893,6 +901,8 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
});
};
const onMessageError = () => {
// Message error may indicate worker is in bad state
terminated = true;
finalize?.({
type: "done",
exitCode: 1,
@@ -902,6 +912,8 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
};
const onClose = () => {
// Worker terminated unexpectedly (crashed or was killed without sending done)
// Mark as terminated since the worker is already dead - calling terminate() again would crash
terminated = true;
const abortMessage =
abortSent && abortReason === "signal"
? "Worker terminated after abort"
@@ -932,11 +944,19 @@ export async function runSubprocess(options: ExecutorOptions): Promise<SingleRes
}
});
// Cleanup
try {
worker.terminate();
} catch {
// Ignore termination errors
// Cleanup - cancel any pending timeouts first
if (terminationTimeoutId) {
clearTimeout(terminationTimeoutId);
terminationTimeoutId = null;
}
cancelPendingTermination();
if (!terminated) {
terminated = true;
try {
worker.terminate();
} catch {
// Ignore termination errors
}
}
let exitCode = done.exitCode;
@@ -15,7 +15,7 @@
import type { AgentEvent, ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, Model } from "@oh-my-pi/pi-ai";
import { logger, untilAborted } from "@oh-my-pi/pi-utils";
import { logger, postmortem, untilAborted } from "@oh-my-pi/pi-utils";
import type { TSchema } from "@sinclair/typebox";
import lspDescription from "../../../prompts/tools/lsp.md" with { type: "text" };
import type { AgentSessionEvent } from "../../agent-session";
@@ -377,17 +377,29 @@ function createPythonProxyTool(): CustomTool<typeof pythonSchema> {
description: getPythonToolDescription(),
parameters: pythonSchema,
execute: async (_toolCallId, params, _onUpdate, _ctx, signal) => {
const timeoutMs = getPythonCallTimeoutMs(params as PythonToolParams);
const result = await callPythonToolViaParent(params as PythonToolParams, signal, timeoutMs);
return {
content:
result?.content?.map((c) =>
c.type === "text"
? { type: "text" as const, text: c.text ?? "" }
: { type: "text" as const, text: JSON.stringify(c) },
) ?? [],
details: result?.details as PythonToolDetails | undefined,
};
try {
const timeoutMs = getPythonCallTimeoutMs(params as PythonToolParams);
const result = await callPythonToolViaParent(params as PythonToolParams, signal, timeoutMs);
return {
content:
result?.content?.map((c) =>
c.type === "text"
? { type: "text" as const, text: c.text ?? "" }
: { type: "text" as const, text: JSON.stringify(c) },
) ?? [],
details: result?.details as PythonToolDetails | undefined,
};
} catch (error) {
return {
content: [
{
type: "text" as const,
text: `Python error: ${error instanceof Error ? error.message : String(error)}`,
},
],
details: { isError: true } as PythonToolDetails,
};
}
},
};
}
@@ -780,7 +792,14 @@ function handleAbort(): void {
}
}
const reportFatal = (message: string): void => {
const reportFatal = async (message: string): Promise<void> => {
// Run postmortem cleanup first to ensure child processes are killed
try {
await postmortem.cleanup();
} catch {
// Ignore cleanup errors
}
const runState = activeRun;
if (runState) {
runState.abortController.abort();
@@ -821,6 +840,16 @@ self.addEventListener("error", (event) => {
self.addEventListener("unhandledrejection", (event) => {
const reason = event.reason;
const message = reason instanceof Error ? reason.stack || reason.message : String(reason);
// Avoid terminating active runs on tool-level errors that bubble as rejections.
if (activeRun) {
logger.error("Unhandled rejection in subagent worker", { error: message });
if ("preventDefault" in event && typeof event.preventDefault === "function") {
event.preventDefault();
}
return;
}
reportFatal(`Unhandled rejection: ${message}`);
});
@@ -4,8 +4,8 @@ import * as path from "node:path";
import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core";
import type { Component } from "@oh-my-pi/pi-tui";
import { Text } from "@oh-my-pi/pi-tui";
import { ptree } from "@oh-my-pi/pi-utils";
import { type Static, Type } from "@sinclair/typebox";
import { $ } from "bun";
import { nanoid } from "nanoid";
import { parse as parseHtml } from "node-html-parser";
import { type Theme, theme } from "../../modes/interactive/theme/theme";
@@ -75,18 +75,58 @@ const CONVERTIBLE_EXTENSIONS = new Set([
* Execute a command and return stdout
*/
type WritableLike = {
write: (chunk: string | Uint8Array) => unknown;
flush?: () => unknown;
end?: () => unknown;
};
const textEncoder = new TextEncoder();
async function writeStdin(handle: unknown, input: string | Buffer): Promise<void> {
if (!handle || typeof handle === "number") return;
if (typeof (handle as WritableStream<Uint8Array>).getWriter === "function") {
const writer = (handle as WritableStream<Uint8Array>).getWriter();
try {
const chunk = typeof input === "string" ? textEncoder.encode(input) : new Uint8Array(input);
await writer.write(chunk);
} finally {
await writer.close();
}
return;
}
const sink = handle as WritableLike;
sink.write(input);
if (sink.flush) sink.flush();
if (sink.end) sink.end();
}
async function exec(
cmd: string,
args: string[],
options?: { timeout?: number; input?: string | Buffer },
): Promise<{ stdout: string; stderr: string; ok: boolean }> {
void options;
const result = await $`${cmd} ${args}`.quiet().nothrow();
const decoder = new TextDecoder();
const proc = ptree.cspawn([cmd, ...args], {
stdin: options?.input ? "pipe" : null,
timeout: options?.timeout,
});
if (options?.input) {
await writeStdin(proc.stdin, options.input);
}
const [stdout, stderr] = await Promise.all([proc.stdout.text(), proc.stderr.text()]);
try {
await proc.exited;
} catch {
// Handle non-zero exit or timeout
}
return {
stdout: result.stdout ? decoder.decode(result.stdout) : "",
stderr: result.stderr ? decoder.decode(result.stderr) : "",
ok: result.exitCode === 0,
stdout,
stderr,
ok: proc.exitCode === 0,
};
}
+7 -5
View File
@@ -92,11 +92,13 @@ async function runInteractiveMode(
await mode.init();
versionCheckPromise.then((newVersion) => {
if (newVersion) {
mode.showNewVersionNotification(newVersion);
}
});
versionCheckPromise
.then((newVersion) => {
if (newVersion) {
mode.showNewVersionNotification(newVersion);
}
})
.catch(() => {});
mode.renderInitialMessages();
@@ -109,7 +109,9 @@ export class LoginDialogComponent extends Container {
showManualInput(prompt: string): Promise<string> {
this.contentContainer.addChild(new Spacer(1));
this.contentContainer.addChild(new Text(theme.fg("dim", prompt), 1, 0));
this.contentContainer.addChild(this.input);
if (!this.contentContainer.children.includes(this.input)) {
this.contentContainer.addChild(this.input);
}
this.contentContainer.addChild(new Text(theme.fg("dim", "(Escape to cancel)"), 1, 0));
this.tui.requestRender();
@@ -129,7 +131,9 @@ export class LoginDialogComponent extends Container {
if (placeholder) {
this.contentContainer.addChild(new Text(theme.fg("dim", `e.g., ${placeholder}`), 1, 0));
}
this.contentContainer.addChild(this.input);
if (!this.contentContainer.children.includes(this.input)) {
this.contentContainer.addChild(this.input);
}
this.contentContainer.addChild(new Text(theme.fg("dim", "(Escape to cancel, Enter to submit)"), 1, 0));
this.input.setValue("");
@@ -260,6 +260,10 @@ export class ToolExecutionComponent extends Container {
): void {
this.result = result;
this.isPartial = isPartial;
// When tool is complete, ensure args are marked complete so spinner stops
if (!isPartial) {
this.argsComplete = true;
}
this.updateSpinnerAnimation();
this.updateDisplay();
// Convert non-PNG images to PNG for Kitty protocol (async)
@@ -168,6 +168,9 @@ export class InteractiveMode implements InteractiveModeContext {
this.editor.onAutocompleteCancel = () => {
this.ui.requestRender(true);
};
this.editor.onAutocompleteUpdate = () => {
this.ui.requestRender(true);
};
try {
this.historyStorage = HistoryStorage.open();
this.editor.setHistoryStorage(this.historyStorage);
@@ -1,6 +1,5 @@
import { afterEach, beforeEach, describe, expect, it } from "bun:test";
import * as fs from "node:fs";
import { tmpdir } from "node:os";
import * as path from "node:path";
import { fileURLToPath } from "node:url";
import { createTempDirSync } from "@oh-my-pi/pi-utils";
@@ -13,7 +12,7 @@ describe("extensions discovery", () => {
let extensionsDir: string;
beforeEach(() => {
tempDir = createTempDirSync(path.join(tmpdir(), "pi-ext-test-"));
tempDir = createTempDirSync("@pi-ext-test-");
extensionsDir = path.join(tempDir.path, ".omp", "extensions");
fs.mkdirSync(extensionsDir, { recursive: true });
});
+20 -6
View File
@@ -5,9 +5,8 @@ export class AbortError extends Error {
assert(signal.aborted, "Abort signal must be aborted");
const message = signal.reason instanceof Error ? signal.reason.message : "Cancelled";
super(`Aborted: ${message}`, { cause: message });
super(`Aborted: ${message}`, { cause: signal.reason });
this.name = "AbortError";
this.cause = signal.reason;
}
}
@@ -42,9 +41,11 @@ export function createAbortablePromise<T>(signal?: AbortSignal): {
reject(new AbortError(signal));
};
signal.addEventListener("abort", abortHandler, { once: true });
promise.finally(() => {
signal.removeEventListener("abort", abortHandler);
});
promise
.finally(() => {
signal.removeEventListener("abort", abortHandler);
})
.catch(() => {});
return { promise, resolve, reject };
}
@@ -63,7 +64,20 @@ export function untilAborted<T>(signal: AbortSignal | undefined | null, pr: () =
return Promise.reject(new AbortError(signal));
}
const { promise, resolve, reject } = createAbortablePromise<T>(signal);
pr().then(resolve, reject);
let settled = false;
const wrappedResolve = (value: T | PromiseLike<T>) => {
if (!settled) {
settled = true;
resolve(value);
}
};
const wrappedReject = (reason?: unknown) => {
if (!settled) {
settled = true;
reject(reason);
}
};
pr().then(wrappedResolve, wrappedReject);
return promise;
}
+58 -28
View File
@@ -6,6 +6,7 @@
* allow reliably releasing resources or shutting down subprocesses, files, sockets, etc.
*/
import { isMainThread } from "node:worker_threads";
import { logger } from ".";
// Cleanup reasons, in order of priority/meaning.
@@ -45,7 +46,8 @@ function runCleanup(reason: Reason): Promise<void> {
// Call .cleanup() for each callback that is still "armed".
// Use Promise.try to handle sync/async, but only those armed.
const promises = callbackList.reverse().map((callback) => {
// Create a copy to avoid mutating the original array with reverse()
const promises = [...callbackList].reverse().map((callback) => {
return Promise.try(() => callback(reason));
});
@@ -61,33 +63,45 @@ function runCleanup(reason: Reason): Promise<void> {
}
// Register signal and error event handlers to trigger cleanup before exit.
process
.on("SIGINT", async () => {
await runCleanup(Reason.SIGINT);
process.exit(130); // 128 + SIGINT (2)
})
.on("uncaughtException", async (err) => {
logger.error("Uncaught exception", { err, stack: err.stack });
await runCleanup(Reason.UNCAUGHT_EXCEPTION);
process.exit(1);
})
.on("unhandledRejection", async (reason) => {
const err = reason instanceof Error ? reason : new Error(String(reason));
logger.error("Unhandled rejection", { err, stack: err.stack });
await runCleanup(Reason.UNHANDLED_REJECTION);
process.exit(1);
})
.on("exit", async () => {
void runCleanup(Reason.EXIT); // fire and forget (exit imminent)
})
.on("SIGTERM", async () => {
await runCleanup(Reason.SIGTERM);
process.exit(143); // 128 + SIGTERM (15)
})
.on("SIGHUP", async () => {
await runCleanup(Reason.SIGHUP);
process.exit(129); // 128 + SIGHUP (1)
// Main thread: full signal handling (SIGINT, SIGTERM, SIGHUP) + exceptions + exit
// Worker thread: exit only (workers use self.addEventListener for exceptions)
if (isMainThread) {
process
.on("SIGINT", async () => {
await runCleanup(Reason.SIGINT);
process.exit(130); // 128 + SIGINT (2)
})
.on("uncaughtException", async (err) => {
logger.error("Uncaught exception", { err, stack: err.stack });
await runCleanup(Reason.UNCAUGHT_EXCEPTION);
process.exit(1);
})
.on("unhandledRejection", async (reason) => {
const err = reason instanceof Error ? reason : new Error(String(reason));
logger.error("Unhandled rejection", { err, stack: err.stack });
await runCleanup(Reason.UNHANDLED_REJECTION);
process.exit(1);
})
.on("exit", async () => {
void runCleanup(Reason.EXIT); // fire and forget (exit imminent)
})
.on("SIGTERM", async () => {
await runCleanup(Reason.SIGTERM);
process.exit(143); // 128 + SIGTERM (15)
})
.on("SIGHUP", async () => {
await runCleanup(Reason.SIGHUP);
process.exit(129); // 128 + SIGHUP (1)
});
} else {
// Worker thread: only register exit handler for cleanup.
// DO NOT register uncaughtException/unhandledRejection handlers here -
// they would swallow errors before the worker's own handlers (self.addEventListener)
// can report failures back to the parent thread.
process.on("exit", () => {
void runCleanup(Reason.EXIT);
});
}
/**
* Register a process cleanup callback, to be run on shutdown, signal, or fatal error.
@@ -134,10 +148,26 @@ export function register(id: string, callback: (reason: Reason) => void | Promis
}
/**
* Runs all cleanup callbacks and exits the process.
* Runs all cleanup callbacks without exiting.
* Use this in workers or when you need to clean up but continue execution.
*/
export function cleanup(): Promise<void> {
return runCleanup(Reason.MANUAL);
}
/**
* Runs all cleanup callbacks and exits.
*
* In main thread: waits for stdout drain, then calls process.exit().
* In workers: runs cleanup only (process.exit would kill entire process).
*/
export async function quit(code: number = 0): Promise<void> {
await runCleanup(Reason.MANUAL);
if (!isMainThread) {
return; // Workers: cleanup done, let worker exit naturally
}
if (process.stdout.writableLength > 0) {
const { promise, resolve } = Promise.withResolvers<void>();
process.stdout.once("drain", resolve);
+9 -2
View File
@@ -201,7 +201,14 @@ export class ChildProcess {
async killAndWait(): Promise<void> {
// Try killing with SIGTERM, then SIGKILL if it doesn't exit within 1 second
this.kill("SIGTERM");
await Promise.race([this.exited, Bun.sleep(1000).then(() => this.kill("SIGKILL"))]);
const exitedOrTimeout = await Promise.race([
this.exited.then(() => "exited" as const),
Bun.sleep(1000).then(() => "timeout" as const),
]);
if (exitedOrTimeout === "timeout") {
this.kill("SIGKILL");
await this.exited.catch(() => {});
}
}
// Output utilities (aliases for easy chaining)
@@ -329,7 +336,7 @@ export class AbortError extends Exception {
*/
export class TimeoutError extends AbortError {
constructor(timeout: number, stderr: string) {
super(new Error(`Process timed out after ${timeout}ms`), stderr);
super(new Error(`Timed out after ${Math.round(timeout / 1000)}s`), stderr);
}
}
+4
View File
@@ -2,6 +2,10 @@
## [Unreleased]
### Fixed
- Fixed viewport tracking after partial renders to prevent autocomplete list artifacts
## [5.6.7] - 2026-01-18
### Added
+2
View File
@@ -290,6 +290,7 @@ export class Editor implements Component, Focusable {
private autocompleteList?: SelectList;
private isAutocompleting: boolean = false;
private autocompletePrefix: string = "";
public onAutocompleteUpdate?: () => void;
// Paste tracking for large pastes
private pastes: Map<number, string> = new Map();
@@ -689,6 +690,7 @@ export class Editor implements Component, Focusable {
// Only pass arrow keys to the list, not Enter/Tab (we handle those directly)
if (matchesKey(data, "up") || matchesKey(data, "down")) {
this.autocompleteList.handleInput(data);
this.onAutocompleteUpdate?.();
return;
}
+3 -1
View File
@@ -1090,7 +1090,9 @@ export class TUI extends Container {
this.terminal.write(buffer);
// Track cursor position for next render
this.cursorRow = finalCursorRow;
// cursorRow represents end-of-content for viewport calculations,
// hardwareCursorRow tracks actual cursor position (may move to cursorPos below)
this.cursorRow = Math.max(0, newLines.length - 1);
this.hardwareCursorRow = finalCursorRow;
// Position hardware cursor for IME