import { spawn, type ChildProcess } from "node:child_process"; import * as fs from "node:fs"; import * as path from "node:path"; import { fileURLToPath } from "node:url"; import type { ExtensionAPI } from "@earendil-works/pi-coding-agent"; import { Type } from "typebox"; // Distinct from the normal `unsloth` provider: subagent mode preserves the user's Pi config. const provider = "unsloth-studio-subagent"; const maxResultCharacters = 100_000; const maxParallelAgents = 4; const cancelGraceMilliseconds = 2_000; const configPath = process.env.UNSLOTH_PI_SUBAGENT_CONFIG || ""; delete process.env.UNSLOTH_PI_SUBAGENT_CONFIG; let config: Record = {}; if (configPath) { try { const parsed = JSON.parse(fs.readFileSync(configPath, "utf8")); if (!parsed || typeof parsed !== "object" || Array.isArray(parsed)) { throw new Error("expected a JSON object"); } config = parsed; } catch (error) { throw new Error(`Could not read Unsloth subagent configuration: ${error}`); } } const model = typeof config.model === "string" ? config.model : ""; const baseUrl = typeof config.baseUrl === "string" ? config.baseUrl : ""; const apiKey = typeof config.apiKey === "string" ? config.apiKey : ""; const approve = config.approve === true; const contextWindow = positiveInt(config.contextWindow, 32768); const maxTokens = positiveInt(config.maxTokens, Math.min(Math.floor(contextWindow / 4), 8192)); let activeAgents = 0; const waitingAgents: Array<() => boolean> = []; function positiveInt(value: unknown, fallback: number): number { const parsed = Number.parseInt(typeof value === "string" ? value : String(value || ""), 10); return Number.isFinite(parsed) && parsed > 0 ? parsed : fallback; } function finalText(message: any): string { if (message?.role !== "assistant" || !Array.isArray(message.content)) return ""; return message.content .filter((part: any) => part?.type === "text" && typeof part.text === "string") .map((part: any) => part.text) .join("\n") .trim(); } function boundedResult(text: string): string { if (text.length <= maxResultCharacters) return text; return `${text.slice(0, maxResultCharacters)}\n\n[Local agent output truncated]`; } function agentSlotRelease(): () => void { let released = false; return () => { if (released) return; released = true; while (waitingAgents.length) { if (waitingAgents.shift()!()) return; } activeAgents -= 1; }; } function acquireAgentSlot(signal: AbortSignal | undefined): Promise<() => void> { if (signal?.aborted) return Promise.reject(new Error("The local Unsloth agent was cancelled.")); if (activeAgents < maxParallelAgents) { activeAgents += 1; return Promise.resolve(agentSlotRelease()); } return new Promise((resolve, reject) => { let waiting = true; const grant = () => { if (!waiting) return false; waiting = false; signal?.removeEventListener("abort", cancel); resolve(agentSlotRelease()); return true; }; const cancel = () => { if (!waiting) return; waiting = false; const index = waitingAgents.indexOf(grant); if (index >= 0) waitingAgents.splice(index, 1); reject(new Error("The local Unsloth agent was cancelled.")); }; waitingAgents.push(grant); signal?.addEventListener("abort", cancel, { once: true }); }); } function piInvocation(args: string[]): { command: string; args: string[] } { const currentScript = process.argv[1]; const bunVirtualScript = currentScript?.startsWith("/$bunfs/root/"); if (currentScript && !bunVirtualScript && fs.existsSync(currentScript)) { return { command: process.execPath, args: [currentScript, ...args] }; } const executable = path.basename(process.execPath).toLowerCase(); if (!/^(node|bun)(\.exe)?$/.test(executable)) return { command: process.execPath, args }; return { command: "pi", args }; } function signalProcessGroup(child: ChildProcess, signal: NodeJS.Signals): void { if (!child.pid) return; try { process.kill(-child.pid, signal); } catch { try { child.kill(signal); } catch { // The process tree already exited. } } } async function stopChildTree(child: ChildProcess): Promise { if (!child.pid) return; if (process.platform === "win32") { await new Promise((resolve) => { const killer = spawn("taskkill", ["/PID", String(child.pid), "/T", "/F"], { shell: false, stdio: "ignore", windowsHide: true, }); killer.once("error", () => { try { child.kill("SIGKILL"); } catch { // The child already exited. } resolve(); }); killer.once("close", (code) => { if (code !== 0) { try { child.kill("SIGKILL"); } catch { // The child already exited. } } resolve(); }); }); return; } signalProcessGroup(child, "SIGTERM"); await new Promise((resolve) => setTimeout(resolve, cancelGraceMilliseconds)); signalProcessGroup(child, "SIGKILL"); } interface LocalAgentResult { task: string; response: string; transcript: any[]; error?: string; } async function runLocalAgent( task: string, cwd: string, signal: AbortSignal | undefined, onProgress: (result: LocalAgentResult) => void, ): Promise { const extension = fileURLToPath(import.meta.url); const args = [ "--mode", "json", "--print", "--no-session", ...(approve ? ["--approve"] : []), "--provider", provider, "--model", model, "--no-extensions", "--extension", extension, `Task: ${task}`, ]; const invocation = piInvocation(args); let output = ""; let stderr = ""; let childError = ""; let aborted = false; const result: LocalAgentResult = { task, response: "", transcript: [] }; const transcriptEntries = new Set(); const appendTranscript = (messages: any[]): boolean => { let changed = false; for (const message of messages) { const entry = JSON.stringify(message); if (transcriptEntries.has(entry)) continue; transcriptEntries.add(entry); result.transcript.push(message); changed = true; } return changed; }; const processLine = (line: string) => { try { const event = JSON.parse(line); if (event.type === "message_end" && event.message && appendTranscript([event.message])) { onProgress(result); } if ( event.type === "turn_end" && Array.isArray(event.toolResults) && event.toolResults.length && appendTranscript(event.toolResults) ) { onProgress(result); } if (event.type !== "message_end") return; const message = event.message; // Pi reports model/API failures as message_end events while still // exiting 0, so the exit status alone cannot surface them. if (message?.stopReason === "error" || message?.stopReason === "aborted") { childError = (typeof message.errorMessage === "string" && message.errorMessage) || `The local Unsloth agent stopped: ${message.stopReason}.`; return; } const response = finalText(message); if (response) { result.response = boundedResult(response); childError = ""; } } catch { // Ignore non-JSON diagnostic lines. The exit status still reports failures. } }; const exitCode = await new Promise((resolve, reject) => { const child = spawn(invocation.command, invocation.args, { cwd, detached: process.platform !== "win32", shell: false, stdio: ["ignore", "pipe", "pipe"], env: { ...process.env, UNSLOTH_PI_SUBAGENT_CHILD: "1", UNSLOTH_PI_SUBAGENT_CONFIG: configPath, }, }); let cleanup: Promise | undefined; const cancel = () => { if (aborted) return; aborted = true; cleanup = stopChildTree(child); }; child.on("error", (error) => { signal?.removeEventListener("abort", cancel); reject(error); }); child.stdout.on("data", (chunk) => { output += chunk.toString(); const lines = output.split("\n"); output = lines.pop() || ""; for (const line of lines) processLine(line); }); child.stderr.on("data", (chunk) => { stderr = (stderr + chunk.toString()).slice(-100_000); }); child.on("close", async (code) => { signal?.removeEventListener("abort", cancel); await cleanup; if (output.trim()) processLine(output); resolve(code ?? 1); }); signal?.addEventListener("abort", cancel, { once: true }); if (signal?.aborted) cancel(); }); if (aborted) throw new Error("The local Unsloth agent was cancelled."); if (exitCode !== 0) { result.error = stderr.trim() || `The local Unsloth agent exited with code ${exitCode}.`; } if (childError) result.error = boundedResult(childError); if (!result.response && !result.error) result.response = "The local agent returned no text."; return result; } export default function unslothSubagent(pi: ExtensionAPI): void { if (!model || !baseUrl || !apiKey || !configPath) { throw new Error("Unsloth subagent configuration is incomplete."); } pi.registerProvider(provider, { name: "Unsloth Studio", baseUrl, apiKey, api: "openai-completions", authHeader: true, models: [ { id: model, name: `${model} via Unsloth`, reasoning: false, input: ["text"], cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, contextWindow, maxTokens, }, ], }); if (process.env.UNSLOTH_PI_SUBAGENT_CHILD === "1") return; pi.registerTool({ name: "unsloth_agent", label: "Unsloth agent", description: "Run local coding agents powered by Unsloth for debugging, implementation, and codebase research. Use task for one agent. To run multiple independent agents, use tasks; up to four run concurrently. The tool returns only after every requested agent finishes.", parameters: Type.Object({ task: Type.Optional( Type.String({ description: "The complete task for one local Unsloth agent." }), ), tasks: Type.Optional( Type.Array(Type.String({ description: "A complete task for one local Unsloth agent." }), { description: "Independent tasks to run concurrently, one local agent per task.", minItems: 2, maxItems: maxParallelAgents, }), ), }), executionMode: "parallel", async execute(_toolCallId, params, signal, onUpdate, ctx) { const singleTask = typeof params.task === "string" && params.task.trim() ? params.task.trim() : ""; const parallelTasks = Array.isArray(params.tasks) ? params.tasks.map((task) => task.trim()).filter(Boolean) : []; if (Boolean(singleTask) === Boolean(parallelTasks.length)) { throw new Error("Provide exactly one of task or tasks."); } if (parallelTasks.length > maxParallelAgents) { throw new Error(`At most ${maxParallelAgents} local agents can run concurrently.`); } if (parallelTasks.length === 1) { throw new Error("Use task for one local agent, or tasks for two to four agents."); } const tasks = singleTask ? [singleTask] : parallelTasks; const results: Array = new Array(tasks.length); let completed = 0; const details = () => ({ provider, model, mode: tasks.length === 1 ? "single" : "parallel", results: results.filter((result): result is LocalAgentResult => Boolean(result)), }); const emitUpdate = () => { onUpdate?.({ content: [ { type: "text", text: `Local agents: ${completed}/${tasks.length} completed`, }, ], details: details(), }); }; await Promise.all( tasks.map(async (task, index) => { let releaseAgentSlot: (() => void) | undefined; try { releaseAgentSlot = await acquireAgentSlot(signal); results[index] = await runLocalAgent(task, ctx.cwd, signal, (partial) => { results[index] = partial; emitUpdate(); }); } catch (error) { results[index] = { task, response: "", transcript: results[index]?.transcript || [], error: String(error), }; } finally { releaseAgentSlot?.(); completed += 1; emitUpdate(); } }), ); if (signal?.aborted) throw new Error("The local Unsloth agent was cancelled."); const completedResults = results.filter( (result): result is LocalAgentResult => Boolean(result), ); const succeeded = completedResults.filter((result) => !result.error).length; const response = completedResults.length === 1 ? completedResults[0].error || completedResults[0].response : [ `Parallel: ${succeeded}/${tasks.length} local agents succeeded`, ...completedResults.map( (result, index) => `\n### Agent ${index + 1}${result.error ? " failed" : ""}\n\n${result.error || result.response}`, ), ].join("\n"); if (succeeded !== completedResults.length) throw new Error(response); return { content: [{ type: "text", text: response }], details: details(), }; }, }); }