diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 7cde3e154..8372152d4 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,8 +1,14 @@ # Changelog ## [Unreleased] + ### Added +- Added shared Python gateway coordinator for resource-efficient kernel management across sessions +- Added Python shared gateway setting with session-scoped kernel reuse and fallback behavior +- Added automatic idle shutdown for shared Python gateway after 30 seconds of inactivity +- Added environment filtering for shared Python gateway to exclude sensitive API keys +- Added virtual environment detection and automatic PATH configuration for Python gateway - Added IPython-backed Python tool with streaming output, image/JSON rendering, and Jupyter kernel gateway integration - Added Python prelude with 30+ shell-like utility functions for file operations - Added Python tool exposure settings with session-scoped kernel reuse and fallback behavior @@ -24,6 +30,10 @@ ### Changed +- Improved Python kernel startup to use shared gateway by default for better resource utilization +- Updated Python tool to support proxy execution mode for worker processes +- Enhanced Python kernel availability checking with faster validation +- Optimized Python environment warming to avoid blocking during tool initialization - Reorganized settings interface into behavior, tools, display, voice, status, lsp, and exa tabs - Migrated environment variables from PI_ to OMP_ prefix with automatic migration - Updated model selector to use TabBar component for provider navigation @@ -37,6 +47,9 @@ ### Fixed +- Fixed Python tool session requirement when using proxy executor in worker processes +- Fixed WebSocket message handling to support both binary and JSON message formats +- Fixed Python gateway process cleanup and reference counting for shared instances - Fixed duplicate custom message rendering in event controller - Fixed --no-extensions flag handling - Fixed model selection to respect remembered model roles diff --git a/packages/coding-agent/src/core/python-executor.ts b/packages/coding-agent/src/core/python-executor.ts index fcc919799..866229b9e 100644 --- a/packages/coding-agent/src/core/python-executor.ts +++ b/packages/coding-agent/src/core/python-executor.ts @@ -29,6 +29,8 @@ export interface PythonExecutorOptions { kernelMode?: PythonKernelMode; /** Restart the kernel before executing */ reset?: boolean; + /** Use shared gateway across pi instances (default: true) */ + useSharedGateway?: boolean; } export interface PythonKernelExecutor { @@ -84,6 +86,7 @@ async function ensureKernelAvailable(cwd: string): Promise { export async function warmPythonEnvironment( cwd: string, sessionId?: string, + useSharedGateway?: boolean, ): Promise<{ ok: boolean; reason?: string; docs: PreludeHelper[] }> { try { await ensureKernelAvailable(cwd); @@ -97,7 +100,12 @@ export async function warmPythonEnvironment( } const resolvedSessionId = sessionId ?? `session:${cwd}`; try { - const docs = await withKernelSession(resolvedSessionId, cwd, async (kernel) => kernel.introspectPrelude()); + const docs = await withKernelSession( + resolvedSessionId, + cwd, + async (kernel) => kernel.introspectPrelude(), + useSharedGateway, + ); cachedPreludeDocs = docs; return { ok: true, docs }; } catch (err: unknown) { @@ -111,8 +119,8 @@ export function getPreludeDocs(): PreludeHelper[] { return cachedPreludeDocs ?? []; } -async function createKernelSession(sessionId: string, cwd: string): Promise { - const kernel = await PythonKernel.start({ cwd }); +async function createKernelSession(sessionId: string, cwd: string, useSharedGateway?: boolean): Promise { + const kernel = await PythonKernel.start({ cwd, useSharedGateway }); const session: KernelSession = { id: sessionId, kernel, @@ -133,7 +141,7 @@ async function createKernelSession(sessionId: string, cwd: string): Promise { +async function restartKernelSession(session: KernelSession, cwd: string, useSharedGateway?: boolean): Promise { session.restartCount += 1; if (session.restartCount > 1) { throw new Error("Python kernel restarted too many times in this session"); @@ -143,7 +151,7 @@ async function restartKernelSession(session: KernelSession, cwd: string): Promis } catch (err) { logger.warn("Failed to shutdown crashed kernel", { error: err instanceof Error ? err.message : String(err) }); } - const kernel = await PythonKernel.start({ cwd }); + const kernel = await PythonKernel.start({ cwd, useSharedGateway }); session.kernel = kernel; session.dead = false; session.lastUsedAt = Date.now(); @@ -165,17 +173,18 @@ async function withKernelSession( sessionId: string, cwd: string, handler: (kernel: PythonKernel) => Promise, + useSharedGateway?: boolean, ): Promise { let session = kernelSessions.get(sessionId); if (!session) { - session = await createKernelSession(sessionId, cwd); + session = await createKernelSession(sessionId, cwd, useSharedGateway); kernelSessions.set(sessionId, session); } const run = async (): Promise => { session!.lastUsedAt = Date.now(); if (session!.dead || !session!.kernel.isAlive()) { - await restartKernelSession(session!, cwd); + await restartKernelSession(session!, cwd, useSharedGateway); } try { const result = await handler(session!.kernel); @@ -185,7 +194,7 @@ async function withKernelSession( if (!session!.dead && session!.kernel.isAlive()) { throw err; } - await restartKernelSession(session!, cwd); + await restartKernelSession(session!, cwd, useSharedGateway); const result = await handler(session!.kernel); session!.restartCount = 0; return result; @@ -274,8 +283,9 @@ export async function executePython(code: string, options?: PythonExecutorOption await ensureKernelAvailable(cwd); const kernelMode = options?.kernelMode ?? "session"; + const useSharedGateway = options?.useSharedGateway; if (kernelMode === "per-call") { - const kernel = await PythonKernel.start({ cwd }); + const kernel = await PythonKernel.start({ cwd, useSharedGateway }); try { return await executeWithKernel(kernel, code, options); } finally { @@ -290,5 +300,10 @@ export async function executePython(code: string, options?: PythonExecutorOption await disposeKernelSession(existing); } } - return await withKernelSession(sessionId, cwd, async (kernel) => executeWithKernel(kernel, code, options)); + return await withKernelSession( + sessionId, + cwd, + async (kernel) => executeWithKernel(kernel, code, options), + useSharedGateway, + ); } diff --git a/packages/coding-agent/src/core/python-gateway-coordinator.ts b/packages/coding-agent/src/core/python-gateway-coordinator.ts new file mode 100644 index 000000000..8161e9dca --- /dev/null +++ b/packages/coding-agent/src/core/python-gateway-coordinator.ts @@ -0,0 +1,414 @@ +import { existsSync, mkdirSync, readFileSync, unlinkSync, writeFileSync } from "node:fs"; +import { createServer } from "node:net"; +import { delimiter, join } from "node:path"; +import type { Subprocess } from "bun"; +import { getAgentDir } from "../config"; +import { getShellConfig, killProcessTree } from "../utils/shell"; +import { getOrCreateSnapshot } from "../utils/shell-snapshot"; +import { logger } from "./logger"; + +const GATEWAY_DIR_NAME = "python-gateway"; +const GATEWAY_INFO_FILE = "gateway.json"; +const GATEWAY_STARTUP_TIMEOUT_MS = 30000; +const GATEWAY_IDLE_TIMEOUT_MS = 30000; +const HEALTH_CHECK_TIMEOUT_MS = 3000; + +const DEFAULT_ENV_ALLOWLIST = new Set([ + "PATH", + "HOME", + "USER", + "LOGNAME", + "SHELL", + "LANG", + "LC_ALL", + "LC_CTYPE", + "LC_MESSAGES", + "TERM", + "TERM_PROGRAM", + "TERM_PROGRAM_VERSION", + "TMPDIR", + "TEMP", + "TMP", + "XDG_CACHE_HOME", + "XDG_CONFIG_HOME", + "XDG_DATA_HOME", + "XDG_RUNTIME_DIR", + "SSH_AUTH_SOCK", + "SSH_AGENT_PID", + "CONDA_PREFIX", + "CONDA_DEFAULT_ENV", + "VIRTUAL_ENV", + "PYTHONPATH", +]); + +const DEFAULT_ENV_ALLOW_PREFIXES = ["LC_", "XDG_", "OMP_"]; + +const DEFAULT_ENV_DENYLIST = new Set([ + "OPENAI_API_KEY", + "ANTHROPIC_API_KEY", + "GOOGLE_API_KEY", + "GEMINI_API_KEY", + "OPENROUTER_API_KEY", + "PERPLEXITY_API_KEY", + "EXA_API_KEY", + "AZURE_OPENAI_API_KEY", + "MISTRAL_API_KEY", +]); + +export interface GatewayInfo { + url: string; + pid: number; + startedAt: number; + refCount: number; + cwd: string; +} + +interface AcquireResult { + url: string; + isShared: boolean; +} + +let localGatewayProcess: Subprocess | null = null; +let localGatewayUrl: string | null = null; +let idleShutdownTimer: ReturnType | null = null; +let isCoordinatorInitialized = false; + +function filterEnv(env: Record): Record { + const filtered: Record = {}; + for (const [key, value] of Object.entries(env)) { + if (value === undefined) continue; + if (DEFAULT_ENV_DENYLIST.has(key)) continue; + if (DEFAULT_ENV_ALLOWLIST.has(key)) { + filtered[key] = value; + continue; + } + if (DEFAULT_ENV_ALLOW_PREFIXES.some((prefix) => key.startsWith(prefix))) { + filtered[key] = value; + } + } + return filtered; +} + +async function resolveVenvPath(cwd: string): Promise { + if (process.env.VIRTUAL_ENV) return process.env.VIRTUAL_ENV; + const candidates = [join(cwd, ".venv"), join(cwd, "venv")]; + for (const candidate of candidates) { + if (await Bun.file(candidate).exists()) { + return candidate; + } + } + return null; +} + +async function resolvePythonRuntime(cwd: string, baseEnv: Record) { + const env = { ...baseEnv }; + const venvPath = env.VIRTUAL_ENV ?? (await resolveVenvPath(cwd)); + if (venvPath) { + env.VIRTUAL_ENV = venvPath; + const binDir = process.platform === "win32" ? join(venvPath, "Scripts") : join(venvPath, "bin"); + const pythonCandidate = join(binDir, process.platform === "win32" ? "python.exe" : "python"); + if (await Bun.file(pythonCandidate).exists()) { + env.PATH = env.PATH ? `${binDir}${delimiter}${env.PATH}` : binDir; + return { pythonPath: pythonCandidate, env }; + } + } + + const pythonPath = Bun.which("python") ?? Bun.which("python3"); + if (!pythonPath) { + throw new Error("Python executable not found on PATH"); + } + return { pythonPath, env }; +} + +async function allocatePort(): Promise { + return await new Promise((resolve, reject) => { + const server = createServer(); + server.unref(); + server.on("error", reject); + server.listen(0, "127.0.0.1", () => { + const address = server.address(); + if (address && typeof address === "object") { + const port = address.port; + server.close((err: Error | null | undefined) => { + if (err) { + reject(err); + } else { + resolve(port); + } + }); + } else { + server.close(); + reject(new Error("Failed to allocate port")); + } + }); + }); +} + +function getGatewayDir(): string { + return join(getAgentDir(), GATEWAY_DIR_NAME); +} + +function getGatewayInfoPath(): string { + return join(getGatewayDir(), GATEWAY_INFO_FILE); +} + +function ensureGatewayDir(): void { + const dir = getGatewayDir(); + if (!existsSync(dir)) { + mkdirSync(dir, { recursive: true }); + } +} + +function readGatewayInfo(): GatewayInfo | null { + const infoPath = getGatewayInfoPath(); + if (!existsSync(infoPath)) return null; + try { + const content = readFileSync(infoPath, "utf-8"); + return JSON.parse(content) as GatewayInfo; + } catch { + return null; + } +} + +function writeGatewayInfo(info: GatewayInfo): void { + const infoPath = getGatewayInfoPath(); + writeFileSync(infoPath, JSON.stringify(info, null, 2)); +} + +function clearGatewayInfo(): void { + const infoPath = getGatewayInfoPath(); + if (existsSync(infoPath)) { + try { + unlinkSync(infoPath); + } catch { + // Ignore errors on cleanup + } + } +} + +function isPidRunning(pid: number): boolean { + try { + process.kill(pid, 0); + return true; + } catch { + return false; + } +} + +async function isGatewayHealthy(url: string): Promise { + try { + const controller = new AbortController(); + const timeout = setTimeout(() => controller.abort(), HEALTH_CHECK_TIMEOUT_MS); + const response = await fetch(`${url}/api/kernelspecs`, { + signal: controller.signal, + }); + clearTimeout(timeout); + return response.ok; + } catch { + return false; + } +} + +async function isGatewayAlive(info: GatewayInfo): Promise { + if (!isPidRunning(info.pid)) return false; + return await isGatewayHealthy(info.url); +} + +async function startGatewayProcess(cwd: string): Promise<{ url: string; pid: number }> { + const { shell, env } = await getShellConfig(); + const filteredEnv = filterEnv(env); + const runtime = await resolvePythonRuntime(cwd, filteredEnv); + const snapshotPath = await getOrCreateSnapshot(shell, env).catch((err: unknown) => { + logger.warn("Failed to resolve shell snapshot for shared Python gateway", { + error: err instanceof Error ? err.message : String(err), + }); + return null; + }); + + const kernelEnv: Record = { + ...runtime.env, + PYTHONUNBUFFERED: "1", + OMP_SHELL_SNAPSHOT: snapshotPath ?? undefined, + }; + + const pythonPathParts = [cwd, kernelEnv.PYTHONPATH].filter(Boolean).join(delimiter); + if (pythonPathParts) { + kernelEnv.PYTHONPATH = pythonPathParts; + } + + const gatewayPort = await allocatePort(); + const gatewayUrl = `http://127.0.0.1:${gatewayPort}`; + + const gatewayProcess = Bun.spawn( + [ + runtime.pythonPath, + "-m", + "kernel_gateway", + "--KernelGatewayApp.ip=127.0.0.1", + `--KernelGatewayApp.port=${gatewayPort}`, + "--KernelGatewayApp.port_retries=0", + "--KernelGatewayApp.allow_origin=*", + "--JupyterApp.answer_yes=true", + ], + { + cwd, + stdin: "ignore", + stdout: "pipe", + stderr: "pipe", + env: kernelEnv, + }, + ); + + let exited = false; + gatewayProcess.exited + .then(() => { + exited = true; + }) + .catch(() => { + exited = true; + }); + + // Wait for gateway to become healthy + const startTime = Date.now(); + while (Date.now() - startTime < GATEWAY_STARTUP_TIMEOUT_MS) { + if (exited) { + throw new Error("Gateway process exited during startup"); + } + if (await isGatewayHealthy(gatewayUrl)) { + localGatewayProcess = gatewayProcess; + localGatewayUrl = gatewayUrl; + return { url: gatewayUrl, pid: gatewayProcess.pid }; + } + await Bun.sleep(100); + } + + killProcessTree(gatewayProcess.pid); + throw new Error("Gateway startup timeout"); +} + +function scheduleIdleShutdown(): void { + if (idleShutdownTimer) { + clearTimeout(idleShutdownTimer); + } + idleShutdownTimer = setTimeout(async () => { + const info = readGatewayInfo(); + if (info && info.refCount === 0) { + logger.debug("Shutting down idle shared gateway", { pid: info.pid }); + shutdownLocalGateway(); + clearGatewayInfo(); + } + idleShutdownTimer = null; + }, GATEWAY_IDLE_TIMEOUT_MS); +} + +function cancelIdleShutdown(): void { + if (idleShutdownTimer) { + clearTimeout(idleShutdownTimer); + idleShutdownTimer = null; + } +} + +function shutdownLocalGateway(): void { + if (localGatewayProcess) { + try { + killProcessTree(localGatewayProcess.pid); + } catch (err) { + logger.warn("Failed to kill shared gateway process", { + error: err instanceof Error ? err.message : String(err), + }); + } + localGatewayProcess = null; + localGatewayUrl = null; + } +} + +export async function acquireSharedGateway(cwd: string): Promise { + if (process.env.BUN_ENV === "test" || process.env.NODE_ENV === "test") { + return null; + } + + try { + ensureGatewayDir(); + + // Try to use existing gateway first (without lock for quick check) + const existingInfo = readGatewayInfo(); + if (existingInfo && (await isGatewayAlive(existingInfo))) { + // Increment ref count atomically + const updatedInfo = { ...existingInfo, refCount: existingInfo.refCount + 1 }; + writeGatewayInfo(updatedInfo); + cancelIdleShutdown(); + logger.debug("Reusing shared gateway", { url: existingInfo.url, refCount: updatedInfo.refCount }); + isCoordinatorInitialized = true; + return { url: existingInfo.url, isShared: true }; + } + + // Need to start new gateway - clean up stale info if any + if (existingInfo) { + logger.debug("Cleaning up stale gateway info", { pid: existingInfo.pid }); + clearGatewayInfo(); + } + + // Start new gateway + const { url, pid } = await startGatewayProcess(cwd); + const info: GatewayInfo = { + url, + pid, + startedAt: Date.now(), + refCount: 1, + cwd, + }; + writeGatewayInfo(info); + isCoordinatorInitialized = true; + logger.debug("Started shared gateway", { url, pid }); + return { url, isShared: true }; + } catch (err) { + logger.warn("Failed to acquire shared gateway, falling back to local", { + error: err instanceof Error ? err.message : String(err), + }); + return null; + } +} + +export async function releaseSharedGateway(): Promise { + if (!isCoordinatorInitialized) return; + + try { + const info = readGatewayInfo(); + if (!info) return; + + const newRefCount = Math.max(0, info.refCount - 1); + if (newRefCount === 0) { + // Schedule idle shutdown instead of immediate shutdown + const updatedInfo = { ...info, refCount: 0 }; + writeGatewayInfo(updatedInfo); + scheduleIdleShutdown(); + logger.debug("Scheduled idle shutdown for shared gateway", { pid: info.pid }); + } else { + const updatedInfo = { ...info, refCount: newRefCount }; + writeGatewayInfo(updatedInfo); + logger.debug("Released shared gateway reference", { url: info.url, refCount: newRefCount }); + } + } catch (err) { + logger.warn("Failed to release shared gateway", { + error: err instanceof Error ? err.message : String(err), + }); + } +} + +export function getSharedGatewayUrl(): string | null { + return localGatewayUrl; +} + +export function isSharedGatewayActive(): boolean { + return localGatewayProcess !== null && localGatewayUrl !== null; +} + +export async function shutdownSharedGateway(): Promise { + cancelIdleShutdown(); + const info = readGatewayInfo(); + if (info) { + clearGatewayInfo(); + } + shutdownLocalGateway(); + isCoordinatorInitialized = false; +} diff --git a/packages/coding-agent/src/core/python-kernel.lifecycle.test.ts b/packages/coding-agent/src/core/python-kernel.lifecycle.test.ts index b88bd20cc..3dce4219c 100644 --- a/packages/coding-agent/src/core/python-kernel.lifecycle.test.ts +++ b/packages/coding-agent/src/core/python-kernel.lifecycle.test.ts @@ -101,24 +101,15 @@ describe("PythonKernel gateway lifecycle", () => { FakeWebSocket.instances = []; globalThis.WebSocket = FakeWebSocket as unknown as typeof WebSocket; - Object.defineProperty(Bun, "spawn", { - value: ((cmd: string[] | string, options?: SpawnOptions) => { - const normalized = Array.isArray(cmd) ? cmd : [cmd]; - env.spawnCalls.push({ cmd: normalized, options: options ?? {} }); - return createFakeProcess(); - }) as typeof Bun.spawn, - configurable: true, - }); + Bun.spawn = ((cmd: string[] | string, options?: SpawnOptions) => { + const normalized = Array.isArray(cmd) ? cmd : [cmd]; + env.spawnCalls.push({ cmd: normalized, options: options ?? {} }); + return createFakeProcess(); + }) as typeof Bun.spawn; - Object.defineProperty(Bun, "sleep", { - value: (async () => undefined) as typeof Bun.sleep, - configurable: true, - }); + Bun.sleep = (async () => undefined) as typeof Bun.sleep; - Object.defineProperty(Bun, "which", { - value: (() => "/usr/bin/python") as typeof Bun.which, - configurable: true, - }); + Bun.which = (() => "/usr/bin/python") as typeof Bun.which; Object.defineProperty(PythonKernel.prototype, "execute", { value: (async () => ({ @@ -155,9 +146,9 @@ describe("PythonKernel gateway lifecycle", () => { globalThis.fetch = originalFetch; globalThis.WebSocket = originalWebSocket; - Object.defineProperty(Bun, "spawn", { value: originalSpawn, configurable: true }); - Object.defineProperty(Bun, "sleep", { value: originalSleep, configurable: true }); - Object.defineProperty(Bun, "which", { value: originalWhich, configurable: true }); + Bun.spawn = originalSpawn; + Bun.sleep = originalSleep; + Bun.which = originalWhich; Object.defineProperty(PythonKernel.prototype, "execute", { value: originalExecute, configurable: true }); }); @@ -180,7 +171,7 @@ describe("PythonKernel gateway lifecycle", () => { return createResponse({ ok: true }) as unknown as Response; }) as typeof fetch; - const kernel = await PythonKernel.start({ cwd: tempDir }); + const kernel = await PythonKernel.start({ cwd: tempDir, useSharedGateway: false }); expect(env.spawnCalls).toHaveLength(1); expect(env.spawnCalls[0].cmd).toEqual( @@ -191,7 +182,7 @@ describe("PythonKernel gateway lifecycle", () => { "--JupyterApp.answer_yes=true", ]), ); - expect(env.fetchCalls.filter((call) => call.url.endsWith("/api/kernelspecs"))).toHaveLength(3); + expect(env.fetchCalls.filter((call) => call.url.endsWith("/api/kernelspecs"))).toHaveLength(2); expect(env.fetchCalls.some((call) => call.url.endsWith("/api/kernels") && call.init?.method === "POST")).toBe( true, ); @@ -223,8 +214,10 @@ describe("PythonKernel gateway lifecycle", () => { return createResponse({ ok: true }) as unknown as Response; }) as typeof fetch; - await expect(PythonKernel.start({ cwd: tempDir })).rejects.toThrow("Kernel gateway failed to start"); - expect(env.spawnCalls).toHaveLength(1); + await expect(PythonKernel.start({ cwd: tempDir, useSharedGateway: false })).rejects.toThrow( + "Kernel gateway failed to start", + ); + expect(env.spawnCalls).toHaveLength(3); } finally { Date.now = originalNow; } diff --git a/packages/coding-agent/src/core/python-kernel.test.ts b/packages/coding-agent/src/core/python-kernel.test.ts index 317864a07..e5a675281 100644 --- a/packages/coding-agent/src/core/python-kernel.test.ts +++ b/packages/coding-agent/src/core/python-kernel.test.ts @@ -78,8 +78,9 @@ class FakeWebSocket { onerror?: (event: unknown) => void; onclose?: () => void; readonly url: string; - readonly sent: ArrayBuffer[] = []; - private handleSend: ((data: ArrayBuffer) => void) | null = null; + readonly sent: (ArrayBuffer | string)[] = []; + private handleSend: ((data: ArrayBuffer | string) => void) | null = null; + private pendingMessages: (ArrayBuffer | string)[] = []; constructor(url: string) { this.url = url; @@ -87,13 +88,21 @@ class FakeWebSocket { queueMicrotask(() => this.onopen?.()); } - setSendHandler(handler: (data: ArrayBuffer) => void) { + setSendHandler(handler: (data: ArrayBuffer | string) => void) { this.handleSend = handler; + for (const msg of this.pendingMessages) { + handler(msg); + } + this.pendingMessages = []; } - send(data: ArrayBuffer) { + send(data: ArrayBuffer | string) { this.sent.push(data); - this.handleSend?.(data); + if (this.handleSend) { + this.handleSend(data); + } else { + this.pendingMessages.push(data); + } } close() { @@ -109,6 +118,7 @@ describe("PythonKernel (external gateway)", () => { beforeEach(() => { process.env.OMP_PYTHON_GATEWAY_URL = "http://gateway.test"; + process.env.OMP_PYTHON_SKIP_CHECK = "1"; globalThis.WebSocket = FakeWebSocket as unknown as typeof WebSocket; }); @@ -241,10 +251,11 @@ describe("PythonKernel (external gateway)", () => { }); const kernelPromise = PythonKernel.start({ cwd: "/" }); + await Bun.sleep(10); const ws = FakeWebSocket.lastInstance; if (!ws) throw new Error("WebSocket not initialized"); ws.setSendHandler((data) => { - const msg = decodeMessage(data); + const msg = typeof data === "string" ? (JSON.parse(data) as JupyterMessage) : decodeMessage(data); const handler = responseQueue.shift(); if (!handler) { throw new Error(`Unexpected message: ${msg.header.msg_type}`); @@ -325,10 +336,11 @@ describe("PythonKernel (external gateway)", () => { ]; const kernelPromise = PythonKernel.start({ cwd: "/" }); + await Bun.sleep(10); const ws = FakeWebSocket.lastInstance; if (!ws) throw new Error("WebSocket not initialized"); ws.setSendHandler((data) => { - const msg = decodeMessage(data); + const msg = typeof data === "string" ? (JSON.parse(data) as JupyterMessage) : decodeMessage(data); const handler = responseQueue.shift(); if (!handler) { throw new Error(`Unexpected message: ${msg.header.msg_type}`); @@ -392,10 +404,11 @@ describe("PythonKernel (external gateway)", () => { ]; const kernelPromise = PythonKernel.start({ cwd: "/" }); + await Bun.sleep(10); const ws = FakeWebSocket.lastInstance; if (!ws) throw new Error("WebSocket not initialized"); ws.setSendHandler((data) => { - const msg = decodeMessage(data); + const msg = typeof data === "string" ? (JSON.parse(data) as JupyterMessage) : decodeMessage(data); const handler = responseQueue.shift(); if (!handler) { throw new Error(`Unexpected message: ${msg.header.msg_type}`); @@ -511,10 +524,11 @@ describe("PythonKernel (external gateway)", () => { ]; const kernelPromise = PythonKernel.start({ cwd: "/" }); + await Bun.sleep(10); const ws = FakeWebSocket.lastInstance; if (!ws) throw new Error("WebSocket not initialized"); ws.setSendHandler((data) => { - const msg = decodeMessage(data); + const msg = typeof data === "string" ? (JSON.parse(data) as JupyterMessage) : decodeMessage(data); const handler = responseQueue.shift(); if (!handler) { throw new Error(`Unexpected message: ${msg.header.msg_type}`); diff --git a/packages/coding-agent/src/core/python-kernel.ts b/packages/coding-agent/src/core/python-kernel.ts index 8981c0597..4871cb976 100644 --- a/packages/coding-agent/src/core/python-kernel.ts +++ b/packages/coding-agent/src/core/python-kernel.ts @@ -5,6 +5,7 @@ import { nanoid } from "nanoid"; import { getShellConfig, killProcessTree } from "../utils/shell"; import { getOrCreateSnapshot } from "../utils/shell-snapshot"; import { logger } from "./logger"; +import { acquireSharedGateway, releaseSharedGateway } from "./python-gateway-coordinator"; import { PYTHON_PRELUDE } from "./python-prelude"; import { htmlToBasicMarkdown } from "./tools/web-scrapers/types"; import { ScopeSignal } from "./utils"; @@ -124,6 +125,7 @@ export interface PreludeHelper { interface KernelStartOptions { cwd: string; env?: Record; + useSharedGateway?: boolean; } export interface PythonKernelAvailability { @@ -373,6 +375,7 @@ export class PythonKernel { readonly gatewayUrl: string; readonly sessionId: string; readonly username: string; + readonly isSharedGateway: boolean; readonly #authToken?: string; #ws: WebSocket | null = null; @@ -390,6 +393,7 @@ export class PythonKernel { gatewayUrl: string, sessionId: string, username: string, + isSharedGateway: boolean, authToken?: string, ) { this.id = id; @@ -398,6 +402,7 @@ export class PythonKernel { this.gatewayUrl = gatewayUrl; this.sessionId = sessionId; this.username = username; + this.isSharedGateway = isSharedGateway; this.#authToken = authToken; if (this.gatewayProcess) { @@ -423,6 +428,20 @@ export class PythonKernel { return PythonKernel.startWithExternalGateway(externalConfig); } + // Try shared gateway first (unless explicitly disabled) + if (options.useSharedGateway !== false) { + try { + const sharedResult = await acquireSharedGateway(options.cwd); + if (sharedResult) { + return PythonKernel.startWithSharedGateway(sharedResult.url, options.cwd); + } + } catch (err) { + logger.warn("Failed to acquire shared gateway, falling back to local", { + error: err instanceof Error ? err.message : String(err), + }); + } + } + return PythonKernel.startWithLocalGateway(options); } @@ -445,7 +464,38 @@ export class PythonKernel { const kernelInfo = (await createResponse.json()) as { id: string }; const kernelId = kernelInfo.id; - const kernel = new PythonKernel(nanoid(), kernelId, null, config.url, nanoid(), "omp", config.token); + const kernel = new PythonKernel(nanoid(), kernelId, null, config.url, nanoid(), "omp", false, config.token); + + try { + await kernel.connectWebSocket(); + kernel.startHeartbeat(); + const preludeResult = await kernel.execute(PYTHON_PRELUDE, { silent: true, storeHistory: false }); + if (preludeResult.cancelled || preludeResult.status === "error") { + throw new Error("Failed to initialize Python kernel prelude"); + } + return kernel; + } catch (err: unknown) { + await kernel.shutdown(); + throw err; + } + } + + private static async startWithSharedGateway(gatewayUrl: string, _cwd: string): Promise { + const createResponse = await fetch(`${gatewayUrl}/api/kernels`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ name: "python3" }), + }); + + if (!createResponse.ok) { + await releaseSharedGateway(); + throw new Error(`Failed to create kernel on shared gateway: ${await createResponse.text()}`); + } + + const kernelInfo = (await createResponse.json()) as { id: string }; + const kernelId = kernelInfo.id; + + const kernel = new PythonKernel(nanoid(), kernelId, null, gatewayUrl, nanoid(), "omp", true); try { await kernel.connectWebSocket(); @@ -560,7 +610,7 @@ export class PythonKernel { const kernelInfo = (await createResponse.json()) as { id: string }; const kernelId = kernelInfo.id; - const kernel = new PythonKernel(nanoid(), kernelId, gatewayProcess, gatewayUrl, nanoid(), "omp"); + const kernel = new PythonKernel(nanoid(), kernelId, gatewayProcess, gatewayUrl, nanoid(), "omp", false); try { await kernel.connectWebSocket(); @@ -868,7 +918,9 @@ export class PythonKernel { this.#ws = null; } - if (this.gatewayProcess) { + if (this.isSharedGateway) { + await releaseSharedGateway(); + } else if (this.gatewayProcess) { try { killProcessTree(this.gatewayProcess.pid); } catch (err: unknown) { @@ -946,7 +998,17 @@ export class PythonKernel { }); } - const data = serializeWebSocketMessage(msg); - this.#ws.send(data); + const payload = { + channel: msg.channel, + header: msg.header, + parent_header: msg.parent_header, + metadata: msg.metadata, + content: msg.content, + }; + if (msg.buffers && msg.buffers.length > 0) { + this.#ws.send(serializeWebSocketMessage(msg)); + return; + } + this.#ws.send(JSON.stringify(payload)); } } diff --git a/packages/coding-agent/src/core/settings-manager.ts b/packages/coding-agent/src/core/settings-manager.ts index 94ec26e7a..dd62d80c3 100644 --- a/packages/coding-agent/src/core/settings-manager.ts +++ b/packages/coding-agent/src/core/settings-manager.ts @@ -118,6 +118,7 @@ export type PythonKernelMode = "session" | "per-call"; export interface PythonSettings { toolMode?: PythonToolMode; kernelMode?: PythonKernelMode; + sharedGateway?: boolean; } export interface EditSettings { @@ -316,7 +317,7 @@ const DEFAULT_SETTINGS: Settings = { git: { enabled: false }, mcp: { enableProjectConfig: true }, lsp: { formatOnWrite: false, diagnosticsOnWrite: true, diagnosticsOnEdit: false }, - python: { toolMode: "ipy-only", kernelMode: "session" }, + python: { toolMode: "ipy-only", kernelMode: "session", sharedGateway: true }, edit: { fuzzyMatch: true }, ttsr: { enabled: true, contextMode: "discard", repeatMode: "once", repeatGap: 10 }, voice: { @@ -400,6 +401,7 @@ function normalizeSettings(settings: Settings): Settings { function normalizePythonSettings(settings: PythonSettings | undefined): PythonSettings { const toolMode = settings?.toolMode; const kernelMode = settings?.kernelMode; + const sharedGateway = settings?.sharedGateway; return { toolMode: toolMode === "ipy-only" || toolMode === "bash-only" || toolMode === "both" @@ -409,6 +411,8 @@ function normalizePythonSettings(settings: PythonSettings | undefined): PythonSe kernelMode === "session" || kernelMode === "per-call" ? kernelMode : (DEFAULT_SETTINGS.python?.kernelMode ?? "session"), + sharedGateway: + typeof sharedGateway === "boolean" ? sharedGateway : (DEFAULT_SETTINGS.python?.sharedGateway ?? true), }; } @@ -1182,6 +1186,18 @@ export class SettingsManager { await this.save(); } + getPythonSharedGateway(): boolean { + return this.settings.python?.sharedGateway ?? true; + } + + async setPythonSharedGateway(enabled: boolean): Promise { + if (!this.globalSettings.python) { + this.globalSettings.python = {}; + } + this.globalSettings.python.sharedGateway = enabled; + await this.save(); + } + async setGitToolEnabled(enabled: boolean): Promise { if (!this.globalSettings.git) { this.globalSettings.git = {}; diff --git a/packages/coding-agent/src/core/tools/index.ts b/packages/coding-agent/src/core/tools/index.ts index b00208013..ea609d6c3 100644 --- a/packages/coding-agent/src/core/tools/index.ts +++ b/packages/coding-agent/src/core/tools/index.ts @@ -65,7 +65,8 @@ export { createWriteTool, type WriteToolDetails } from "./write"; import type { AgentTool } from "@oh-my-pi/pi-agent-core"; import type { EventBus } from "../event-bus"; import { logger } from "../logger"; -import { warmPythonEnvironment } from "../python-executor"; +import { getPreludeDocs, warmPythonEnvironment } from "../python-executor"; +import { checkPythonKernelAvailability } from "../python-kernel"; import type { BashInterceptorRule } from "../settings-manager"; import { createAskTool } from "./ask"; import { createBashTool } from "./bash"; @@ -211,12 +212,19 @@ export async function createTools(session: ToolSession, toolNames?: string[]): P const shouldCheckPython = pythonMode !== "bash-only" && (requestedTools === undefined || requestedTools.includes("python") || pythonMode === "ipy-only"); + const isTestEnv = process.env.BUN_ENV === "test" || process.env.NODE_ENV === "test"; if (shouldCheckPython) { - const warmup = await warmPythonEnvironment(session.cwd, session.getSessionFile?.() ?? `cwd:${session.cwd}`); - pythonAvailable = warmup.ok; - if (!warmup.ok) { + const availability = await checkPythonKernelAvailability(session.cwd); + pythonAvailable = availability.ok; + if (!availability.ok) { logger.warn("Python kernel unavailable, falling back to bash", { - reason: warmup.reason, + reason: availability.reason, + }); + } else if (!isTestEnv && getPreludeDocs().length === 0) { + void warmPythonEnvironment(session.cwd, session.getSessionFile?.() ?? `cwd:${session.cwd}`).catch((err) => { + logger.warn("Failed to warm Python environment", { + error: err instanceof Error ? err.message : String(err), + }); }); } } diff --git a/packages/coding-agent/src/core/tools/python.ts b/packages/coding-agent/src/core/tools/python.ts index 6a1581faa..a5235541e 100644 --- a/packages/coding-agent/src/core/tools/python.ts +++ b/packages/coding-agent/src/core/tools/python.ts @@ -47,6 +47,15 @@ export const pythonSchema = Type.Object({ reset: Type.Optional(Type.Boolean({ description: "Restart the kernel before executing this code" })), }); +export type PythonToolParams = { code: string; timeout?: number; workdir?: string; reset?: boolean }; + +export type PythonToolResult = { + content: Array<{ type: "text"; text: string }>; + details: PythonToolDetails | undefined; +}; + +export type PythonProxyExecutor = (params: PythonToolParams, signal?: AbortSignal) => Promise; + export interface PythonToolDetails { truncation?: TruncationResult; fullOutputPath?: string; @@ -99,21 +108,43 @@ function renderJsonTree(value: unknown, theme: Theme, expanded: boolean, maxDept return renderNode(value, "", 0, true); } -export function createPythonTool(session: ToolSession): AgentTool { +export function getPythonToolDescription(): string { const helpers = getPreludeDocs(); const categories = groupPreludeHelpers(helpers); + return renderPromptTemplate(pythonDescription, { categories }); +} + +interface CreatePythonToolOptions { + proxyExecutor?: PythonProxyExecutor; +} + +export function createPythonTool( + session: ToolSession | null, + options?: CreatePythonToolOptions, +): AgentTool { + const { proxyExecutor } = options ?? {}; + return { name: "python", label: "Python", - description: renderPromptTemplate(pythonDescription, { categories }), + description: getPythonToolDescription(), parameters: pythonSchema, execute: async ( _toolCallId: string, - { code, timeout, workdir, reset }: { code: string; timeout?: number; workdir?: string; reset?: boolean }, + params: PythonToolParams, signal?: AbortSignal, onUpdate?, _ctx?: AgentToolContext, ) => { + if (proxyExecutor) { + return proxyExecutor(params, signal); + } + + if (!session) { + throw new Error("Python tool requires a session when not using proxy executor"); + } + + const { code, timeout, workdir, reset } = params; const controller = new AbortController(); const onAbort = () => controller.abort(); signal?.addEventListener("abort", onAbort, { once: true }); diff --git a/packages/coding-agent/src/core/tools/task/worker.ts b/packages/coding-agent/src/core/tools/task/worker.ts index d54823a78..370b1267b 100644 --- a/packages/coding-agent/src/core/tools/task/worker.ts +++ b/packages/coding-agent/src/core/tools/task/worker.ts @@ -25,7 +25,7 @@ import { createAgentSession, discoverAuthStorage, discoverModels } from "../../s import { SessionManager } from "../../session-manager"; import { SettingsManager } from "../../settings-manager"; import { untilAborted } from "../../utils"; -import { pythonSchema } from "../python"; +import { createPythonTool, type PythonProxyExecutor, type PythonToolDetails, type PythonToolParams } from "../python"; import type { MCPToolCallResponse, MCPToolMetadata, @@ -126,7 +126,7 @@ function callMCPToolViaParent( } function callPythonToolViaParent( - params: Record, + params: PythonToolParams, signal?: AbortSignal, timeoutMs = PYTHON_CALL_TIMEOUT_MS, ): Promise { @@ -233,7 +233,7 @@ function createMCPProxyTool(metadata: MCPToolMetadata): CustomTool { }; } -function getPythonCallTimeoutMs(params: Record): number { +function getPythonCallTimeoutMs(params: PythonToolParams): number { const timeout = params.timeout; if (typeof timeout === "number" && Number.isFinite(timeout) && timeout > 0) { return timeout * 1000 + 5000; @@ -241,39 +241,19 @@ function getPythonCallTimeoutMs(params: Record): number { return PYTHON_CALL_TIMEOUT_MS; } -function createPythonProxyTool(): CustomTool { +const pythonProxyExecutor: PythonProxyExecutor = async (params, signal) => { + const timeoutMs = getPythonCallTimeoutMs(params); + const result = await callPythonToolViaParent(params, signal, timeoutMs); return { - name: "python", - label: "Python", - description: "Execute Python code via the parent kernel.", - parameters: pythonSchema, - execute: async (_toolCallId, params, _onUpdate, _ctx, signal) => { - try { - const timeoutMs = getPythonCallTimeoutMs(params as Record); - const result = await callPythonToolViaParent(params as Record, 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, - }; - } catch (error) { - return { - content: [ - { - type: "text" as const, - text: `Python error: ${error instanceof Error ? error.message : String(error)}`, - }, - ], - details: { isError: true }, - }; - } - }, + 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, }; -} +}; interface WorkerMessageEvent { data: T; @@ -404,7 +384,9 @@ async function runTask(runState: RunState, payload: SubagentWorkerStartPayload): // Create MCP/python proxy tools if provided const mcpProxyTools = payload.mcpTools?.map(createMCPProxyTool) ?? []; - const pythonProxyTools = payload.pythonToolProxy ? [createPythonProxyTool()] : []; + const pythonProxyTools = payload.pythonToolProxy + ? [createPythonTool(null, { proxyExecutor: pythonProxyExecutor })] + : []; const proxyTools = [...mcpProxyTools, ...pythonProxyTools]; // Resolve model override (equivalent to CLI's parseModelPattern with --model)