From 8e624225412d137cfbed2f6c76d4cb26e18f4a17 Mon Sep 17 00:00:00 2001 From: thossullivan Date: Sat, 29 Aug 2026 09:52:58 -0500 Subject: [PATCH 1/2] fix(app-server): unsubscribe task threads after client disconnect --- plugins/codex/scripts/app-server-broker.mjs | 232 ++++++++++- .../scripts/lib/app-server-protocol.d.ts | 3 + tests/broker-subscriptions.test.mjs | 380 ++++++++++++++++++ tests/fake-codex-fixture.mjs | 59 ++- 4 files changed, 672 insertions(+), 2 deletions(-) create mode 100644 tests/broker-subscriptions.test.mjs diff --git a/plugins/codex/scripts/app-server-broker.mjs b/plugins/codex/scripts/app-server-broker.mjs index 1954274fe..0f420e91f 100644 --- a/plugins/codex/scripts/app-server-broker.mjs +++ b/plugins/codex/scripts/app-server-broker.mjs @@ -10,6 +10,48 @@ import { BROKER_BUSY_RPC_CODE, CodexAppServerClient } from "./lib/app-server.mjs import { parseBrokerEndpoint } from "./lib/broker-endpoint.mjs"; const STREAMING_METHODS = new Set(["turn/start", "review/start", "thread/compact/start"]); +const SUBSCRIBING_METHODS = new Set(["thread/start", "thread/resume", "thread/fork"]); + +function buildSubscriptionThreadIds(method, result) { + const threadIds = new Set(); + if (SUBSCRIBING_METHODS.has(method) && result?.thread?.id) { + threadIds.add(result.thread.id); + } + if (method === "review/start" && result?.reviewThreadId) { + threadIds.add(result.reviewThreadId); + } + return threadIds; +} + +function buildProvisionalSubscriptionThreadIds(method, params) { + const threadIds = new Set(); + if (method === "thread/resume" && params?.threadId) { + threadIds.add(params.threadId); + } + if (method === "review/start" && params?.threadId) { + threadIds.add(params.threadId); + } + return threadIds; +} + +function buildNotificationSubscriptionThreadIds(message) { + const threadIds = new Set(); + const params = message?.params; + if (message?.method === "thread/started" && params?.thread?.id) { + threadIds.add(params.thread.id); + } + if (params?.threadId) { + threadIds.add(params.threadId); + } + if (Array.isArray(params?.item?.receiverThreadIds)) { + for (const threadId of params.item.receiverThreadIds) { + if (threadId) { + threadIds.add(threadId); + } + } + } + return threadIds; +} function buildStreamThreadIds(method, params, result) { const threadIds = new Set(); @@ -70,6 +112,175 @@ async function main() { let activeStreamSocket = null; let activeStreamThreadIds = null; const sockets = new Set(); + // App-server subscriptions belong to the broker's single upstream connection. + // Mirror downstream ownership so one client cannot release another client's thread. + const socketThreadIds = new Map(); + const threadSockets = new Map(); + const pendingUnsubscribes = new Map(); + const socketUnsubscribingThreadIds = new Map(); + + function addThreadOwner(socket, threadId) { + let ownedThreadIds = socketThreadIds.get(socket); + if (!ownedThreadIds) { + ownedThreadIds = new Set(); + socketThreadIds.set(socket, ownedThreadIds); + } + if (ownedThreadIds.has(threadId)) { + return false; + } + ownedThreadIds.add(threadId); + + let owners = threadSockets.get(threadId); + if (!owners) { + owners = new Set(); + threadSockets.set(threadId, owners); + } + owners.add(socket); + return true; + } + + function removeThreadOwner(socket, threadId) { + const ownedThreadIds = socketThreadIds.get(socket); + if (!ownedThreadIds?.delete(threadId)) { + return false; + } + if (ownedThreadIds.size === 0) { + socketThreadIds.delete(socket); + } + + const owners = threadSockets.get(threadId); + owners?.delete(socket); + if (owners?.size === 0) { + threadSockets.delete(threadId); + } + return true; + } + + function setSocketThreadUnsubscribing(socket, threadId, isUnsubscribing) { + let threadIds = socketUnsubscribingThreadIds.get(socket); + if (isUnsubscribing) { + if (!threadIds) { + threadIds = new Set(); + socketUnsubscribingThreadIds.set(socket, threadIds); + } + threadIds.add(threadId); + return; + } + threadIds?.delete(threadId); + if (threadIds?.size === 0) { + socketUnsubscribingThreadIds.delete(socket); + } + } + + function requestThreadUnsubscribe(threadId) { + const pending = pendingUnsubscribes.get(threadId); + if (pending) { + return pending; + } + const request = appClient.request("thread/unsubscribe", { threadId }).then( + (result) => ({ result, error: null }), + (error) => { + process.stderr.write( + `Failed to unsubscribe Codex thread ${threadId}: ${error instanceof Error ? error.message : String(error)}\n` + ); + return { result: null, error }; + } + ); + pendingUnsubscribes.set(threadId, request); + void request.finally(() => { + if (pendingUnsubscribes.get(threadId) === request) { + pendingUnsubscribes.delete(threadId); + } + }); + return request; + } + + async function unsubscribeIfUnowned(threadId) { + if (threadSockets.has(threadId) || appClient.closed) { + return null; + } + return requestThreadUnsubscribe(threadId); + } + + async function releaseThreadOwners(socket, threadIds = socketThreadIds.get(socket) ?? new Set()) { + const releasedThreadIds = []; + for (const threadId of [...threadIds]) { + if (removeThreadOwner(socket, threadId) && !threadSockets.has(threadId)) { + releasedThreadIds.push(threadId); + } + } + await Promise.all(releasedThreadIds.map((threadId) => unsubscribeIfUnowned(threadId))); + } + + function trackSubscriptionResults(socket, method, result, provisionalThreadIds) { + for (const threadId of buildSubscriptionThreadIds(method, result)) { + if (provisionalThreadIds.has(threadId)) { + continue; + } + if (socket.destroyed || !sockets.has(socket)) { + void unsubscribeIfUnowned(threadId); + continue; + } + addThreadOwner(socket, threadId); + } + } + + function trackNotificationSubscriptions(socket, message) { + // App-server auto-subscribes its connection to child threads created by subagents. + // Attribute those notification-only subscriptions to the active downstream client. + for (const threadId of buildNotificationSubscriptionThreadIds(message)) { + if (socket && !socket.destroyed && sockets.has(socket)) { + if (!socketUnsubscribingThreadIds.get(socket)?.has(threadId)) { + addThreadOwner(socket, threadId); + } + } else { + void unsubscribeIfUnowned(threadId); + } + } + } + + async function handleThreadUnsubscribe(socket, params) { + const threadId = params?.threadId; + if (typeof threadId !== "string") { + return appClient.request("thread/unsubscribe", params ?? {}); + } + + const ownedThreadIds = socketThreadIds.get(socket); + if (!ownedThreadIds?.has(threadId)) { + if (threadSockets.has(threadId)) { + return { status: "notSubscribed" }; + } + setSocketThreadUnsubscribing(socket, threadId, true); + try { + const outcome = await requestThreadUnsubscribe(threadId); + if (outcome.error) { + throw outcome.error; + } + return outcome.result; + } finally { + setSocketThreadUnsubscribing(socket, threadId, false); + } + } + + removeThreadOwner(socket, threadId); + if (threadSockets.has(threadId)) { + return { status: "unsubscribed" }; + } + + setSocketThreadUnsubscribing(socket, threadId, true); + try { + const outcome = await unsubscribeIfUnowned(threadId); + if (outcome?.error) { + if (!socket.destroyed && sockets.has(socket)) { + addThreadOwner(socket, threadId); + } + throw outcome.error; + } + return outcome?.result ?? { status: "unsubscribed" }; + } finally { + setSocketThreadUnsubscribing(socket, threadId, false); + } + } function clearSocketOwnership(socket) { if (activeRequestSocket === socket) { @@ -83,6 +294,7 @@ async function main() { function routeNotification(message) { const target = activeRequestSocket ?? activeStreamSocket; + trackNotificationSubscriptions(target, message); if (!target) { return; } @@ -195,10 +407,23 @@ async function main() { } const isStreaming = STREAMING_METHODS.has(message.method); + // Claim known thread ids before awaiting app-server. This prevents another + // client's close handler from unsubscribing a concurrently resumed thread. + const provisionalThreadIds = buildProvisionalSubscriptionThreadIds(message.method, message.params ?? {}); + const addedProvisionalThreadIds = new Set(); + for (const threadId of provisionalThreadIds) { + if (addThreadOwner(socket, threadId)) { + addedProvisionalThreadIds.add(threadId); + } + } activeRequestSocket = socket; try { - const result = await appClient.request(message.method, message.params ?? {}); + const result = + message.method === "thread/unsubscribe" + ? await handleThreadUnsubscribe(socket, message.params ?? {}) + : await appClient.request(message.method, message.params ?? {}); + trackSubscriptionResults(socket, message.method, result, provisionalThreadIds); send(socket, { id: message.id, result }); if (isStreaming) { activeStreamSocket = socket; @@ -208,6 +433,7 @@ async function main() { activeRequestSocket = null; } } catch (error) { + await releaseThreadOwners(socket, addedProvisionalThreadIds); send(socket, { id: message.id, error: buildJsonRpcError(error.rpcCode ?? -32000, error.message) @@ -225,11 +451,15 @@ async function main() { socket.on("close", () => { sockets.delete(socket); clearSocketOwnership(socket); + socketUnsubscribingThreadIds.delete(socket); + void releaseThreadOwners(socket); }); socket.on("error", () => { sockets.delete(socket); clearSocketOwnership(socket); + socketUnsubscribingThreadIds.delete(socket); + void releaseThreadOwners(socket); }); }); diff --git a/plugins/codex/scripts/lib/app-server-protocol.d.ts b/plugins/codex/scripts/lib/app-server-protocol.d.ts index f61a4588e..e023b324f 100644 --- a/plugins/codex/scripts/lib/app-server-protocol.d.ts +++ b/plugins/codex/scripts/lib/app-server-protocol.d.ts @@ -21,6 +21,8 @@ import type { ThreadSetNameResponse, ThreadStartParams as RawThreadStartParams, ThreadStartResponse, + ThreadUnsubscribeParams, + ThreadUnsubscribeResponse, Turn, TurnInterruptParams, TurnInterruptResponse, @@ -63,6 +65,7 @@ export interface AppServerMethodMap { "thread/resume": { params: ThreadResumeParams; result: ThreadResumeResponse }; "thread/name/set": { params: ThreadSetNameParams; result: ThreadSetNameResponse }; "thread/list": { params: ThreadListParams; result: ThreadListResponse }; + "thread/unsubscribe": { params: ThreadUnsubscribeParams; result: ThreadUnsubscribeResponse }; "review/start": { params: ReviewStartParams; result: ReviewStartResponse }; "turn/start": { params: TurnStartParams; result: TurnStartResponse }; "turn/interrupt": { params: TurnInterruptParams; result: TurnInterruptResponse }; diff --git a/tests/broker-subscriptions.test.mjs b/tests/broker-subscriptions.test.mjs new file mode 100644 index 000000000..3277b611d --- /dev/null +++ b/tests/broker-subscriptions.test.mjs @@ -0,0 +1,380 @@ +import assert from "node:assert/strict"; +import fs from "node:fs"; +import net from "node:net"; +import path from "node:path"; +import test from "node:test"; +import { spawn } from "node:child_process"; +import { fileURLToPath } from "node:url"; + +import { buildEnv, installFakeCodex } from "./fake-codex-fixture.mjs"; +import { makeTempDir } from "./helpers.mjs"; + +const ROOT = path.resolve(fileURLToPath(new URL("..", import.meta.url))); +const BROKER = path.join(ROOT, "plugins", "codex", "scripts", "app-server-broker.mjs"); + +function delay(ms) { + return new Promise((resolve) => setTimeout(resolve, ms)); +} + +async function waitFor(predicate, { timeoutMs = 10000, intervalMs = 25 } = {}) { + const startedAt = Date.now(); + while (Date.now() - startedAt < timeoutMs) { + if (await predicate()) { + return true; + } + await delay(intervalMs); + } + return false; +} + +function readState(statePath) { + try { + return JSON.parse(fs.readFileSync(statePath, "utf8")); + } catch { + return null; + } +} + +function startBroker(behavior = "review-ok") { + const binDir = makeTempDir("codex-broker-bin-"); + installFakeCodex(binDir, behavior); + const sessionDir = makeTempDir("codex-broker-subscriptions-"); + const cwd = makeTempDir("codex-broker-cwd-"); + const socketPath = path.join(sessionDir, "broker.sock"); + const pidFile = path.join(sessionDir, "broker.pid"); + const statePath = path.join(binDir, "fake-codex-state.json"); + + const child = spawn( + process.execPath, + [BROKER, "serve", "--endpoint", `unix:${socketPath}`, "--cwd", cwd, "--pid-file", pidFile], + { env: buildEnv(binDir), stdio: ["ignore", "pipe", "pipe"] } + ); + let stderr = ""; + child.stderr.on("data", (chunk) => { + stderr += chunk; + }); + const exited = new Promise((resolve) => { + child.on("exit", (code, signal) => resolve({ code, signal })); + }); + + async function stop() { + if (child.exitCode !== null || child.signalCode !== null) { + await exited; + return; + } + child.kill("SIGTERM"); + const result = await Promise.race([exited, delay(5000).then(() => null)]); + if (!result) { + child.kill("SIGKILL"); + await exited; + } + } + + return { + socketPath, + statePath, + stderr: () => stderr, + listening: () => waitFor(() => fs.existsSync(socketPath)), + stop + }; +} + +async function connectClient(socketPath) { + const socket = await new Promise((resolve, reject) => { + const candidate = net.createConnection({ path: socketPath }); + candidate.on("connect", () => resolve(candidate)); + candidate.on("error", reject); + }); + socket.setEncoding("utf8"); + + let nextId = 1; + let buffer = ""; + const pending = new Map(); + const notifications = []; + const notificationWaiters = new Set(); + const closed = new Promise((resolve) => socket.on("close", resolve)); + + function notifyWaiters(message) { + for (const waiter of [...notificationWaiters]) { + if (waiter.predicate(message)) { + notificationWaiters.delete(waiter); + waiter.resolve(message); + } + } + } + + socket.on("data", (chunk) => { + buffer += chunk; + let newlineIndex = buffer.indexOf("\n"); + while (newlineIndex !== -1) { + const line = buffer.slice(0, newlineIndex); + buffer = buffer.slice(newlineIndex + 1); + newlineIndex = buffer.indexOf("\n"); + if (!line.trim()) { + continue; + } + const message = JSON.parse(line); + if (message.id !== undefined) { + const request = pending.get(message.id); + if (request) { + pending.delete(message.id); + if (message.error) { + request.reject(new Error(message.error.message)); + } else { + request.resolve(message.result); + } + } + continue; + } + notifications.push(message); + notifyWaiters(message); + } + }); + + socket.on("error", (error) => { + for (const request of pending.values()) { + request.reject(error); + } + pending.clear(); + }); + + function request(method, params = {}) { + const id = nextId++; + const response = new Promise((resolve, reject) => { + pending.set(id, { resolve, reject }); + }); + socket.write(`${JSON.stringify({ id, method, params })}\n`); + return response; + } + + async function waitForNotification(predicate, timeoutMs = 10000) { + const existing = notifications.find(predicate); + if (existing) { + return existing; + } + const notification = new Promise((resolve) => { + notificationWaiters.add({ predicate, resolve }); + }); + return Promise.race([notification, delay(timeoutMs).then(() => null)]); + } + + await request("initialize", {}); + return { + request, + waitForNotification, + async end() { + socket.end(); + await closed; + }, + destroy() { + socket.destroy(); + } + }; +} + +async function waitForUnsubscribes(statePath, expectedThreadIds) { + const expected = [...expectedThreadIds].sort(); + const found = await waitFor(() => { + const actual = [...(readState(statePath)?.unsubscribeRequests ?? [])].sort(); + return actual.length === expected.length && actual.every((threadId, index) => threadId === expected[index]); + }); + assert.equal(found, true, `expected unsubscribe requests for ${expected.join(", ")}`); +} + +test("broker unsubscribes a completed task thread when its client closes", async (t) => { + const broker = startBroker(); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const started = await client.request("thread/start", { cwd: process.cwd(), ephemeral: true }); + const threadId = started.thread.id; + await client.request("turn/start", { + threadId, + input: [{ type: "text", text: "test normal completion" }] + }); + const completed = await client.waitForNotification( + (message) => message.method === "turn/completed" && message.params?.threadId === threadId + ); + assert.ok(completed, "task never completed"); + + await client.end(); + await waitForUnsubscribes(broker.statePath, [threadId]); + assert.deepEqual(readState(broker.statePath).subscriptions, []); +}); + +test("broker keeps a resumed thread subscribed until its final client closes", async (t) => { + const broker = startBroker(); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const firstClient = await connectClient(broker.socketPath); + const threadId = (await firstClient.request("thread/start", { cwd: process.cwd(), ephemeral: false })).thread.id; + const secondClient = await connectClient(broker.socketPath); + await secondClient.request("thread/resume", { threadId }); + + await firstClient.end(); + await delay(250); + assert.deepEqual(readState(broker.statePath).unsubscribeRequests, []); + + await secondClient.end(); + await waitForUnsubscribes(broker.statePath, [threadId]); +}); + +test("broker unsubscribes source and detached review threads", async (t) => { + const broker = startBroker(); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const sourceThreadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: true })).thread.id; + const review = await client.request("review/start", { + threadId: sourceThreadId, + delivery: "detached", + target: { type: "uncommittedChanges" } + }); + assert.notEqual(review.reviewThreadId, sourceThreadId); + const completed = await client.waitForNotification( + (message) => message.method === "turn/completed" && message.params?.threadId === review.reviewThreadId + ); + assert.ok(completed, "review never completed"); + + await client.end(); + await waitForUnsubscribes(broker.statePath, [sourceThreadId, review.reviewThreadId]); +}); + +test("broker unsubscribes a forked thread when its client closes", async (t) => { + const broker = startBroker(); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const sourceThreadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: false })).thread.id; + const forkThreadId = (await client.request("thread/fork", { threadId: sourceThreadId, ephemeral: true })).thread.id; + + await client.end(); + await waitForUnsubscribes(broker.statePath, [sourceThreadId, forkThreadId]); +}); + +test("broker unsubscribes when a client disconnects during an active turn", async (t) => { + const broker = startBroker("slow-task"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const threadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: true })).thread.id; + await client.request("turn/start", { + threadId, + input: [{ type: "text", text: "disconnect this client" }] + }); + client.destroy(); + + await waitForUnsubscribes(broker.statePath, [threadId]); +}); + +test("broker unsubscribes auto-subscribed subagent threads", async (t) => { + const broker = startBroker("with-subagent"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const threadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: true })).thread.id; + await client.request("turn/start", { + threadId, + input: [{ type: "text", text: "delegate this task" }] + }); + const completed = await client.waitForNotification( + (message) => message.method === "turn/completed" && message.params?.threadId === threadId + ); + assert.ok(completed, "task never completed"); + + const subagentThread = readState(broker.statePath).threads.find((thread) => thread.name === "design-challenger"); + assert.ok(subagentThread, "subagent thread was not created"); + + await client.end(); + await waitForUnsubscribes(broker.statePath, [threadId, subagentThread.id]); + assert.deepEqual(readState(broker.statePath).subscriptions, []); +}); + +test("broker unsubscribes a child thread created after its client disconnects", async (t) => { + const broker = startBroker("with-delayed-subagent"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const threadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: true })).thread.id; + await client.request("turn/start", { + threadId, + input: [{ type: "text", text: "disconnect before delegation" }] + }); + client.destroy(); + + const childCreated = await waitFor(() => + readState(broker.statePath)?.threads.some((thread) => thread.name === "delayed-design-challenger") + ); + assert.equal(childCreated, true, "delayed child thread was not created"); + const childThread = readState(broker.statePath).threads.find( + (thread) => thread.name === "delayed-design-challenger" + ); + + await waitForUnsubscribes(broker.statePath, [threadId, childThread.id]); + assert.deepEqual(readState(broker.statePath).subscriptions, []); +}); + +test("broker keeps shared upstream subscriptions when one client explicitly unsubscribes", async (t) => { + const broker = startBroker(); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const firstClient = await connectClient(broker.socketPath); + const threadId = (await firstClient.request("thread/start", { cwd: process.cwd(), ephemeral: false })).thread.id; + const secondClient = await connectClient(broker.socketPath); + await secondClient.request("thread/resume", { threadId }); + + assert.deepEqual(await firstClient.request("thread/unsubscribe", { threadId }), { status: "unsubscribed" }); + assert.deepEqual(readState(broker.statePath).unsubscribeRequests, []); + assert.deepEqual(readState(broker.statePath).subscriptions, [threadId]); + + const thirdClient = await connectClient(broker.socketPath); + assert.deepEqual(await thirdClient.request("thread/unsubscribe", { threadId }), { status: "notSubscribed" }); + assert.deepEqual(readState(broker.statePath).unsubscribeRequests, []); + + await firstClient.end(); + await thirdClient.end(); + await secondClient.end(); + await waitForUnsubscribes(broker.statePath, [threadId]); + assert.deepEqual(readState(broker.statePath).subscriptions, []); +}); + +test("broker does not reclaim ownership from notifications during explicit unsubscribe", async (t) => { + const broker = startBroker("unsubscribe-notifies"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const threadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: true })).thread.id; + assert.deepEqual(await client.request("thread/unsubscribe", { threadId }), { status: "unsubscribed" }); + await client.end(); + await delay(250); + + assert.deepEqual(readState(broker.statePath).unsubscribeRequests, [threadId]); + assert.deepEqual(readState(broker.statePath).subscriptions, []); +}); + +test("broker logs upstream unsubscribe failures", async (t) => { + const broker = startBroker("unsubscribe-fails"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const threadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: true })).thread.id; + await client.end(); + + const requestObserved = await waitFor(() => readState(broker.statePath)?.unsubscribeRequests?.includes(threadId)); + assert.equal(requestObserved, true, "unsubscribe request was not observed"); + const warningObserved = await waitFor( + () => broker.stderr().includes(`Failed to unsubscribe Codex thread ${threadId}: thread unsubscribe failed`) + ); + assert.equal(warningObserved, true, "unsubscribe failure was not logged"); + assert.deepEqual(readState(broker.statePath).subscriptions, [threadId]); +}); diff --git a/tests/fake-codex-fixture.mjs b/tests/fake-codex-fixture.mjs index f83c96a0d..db0efc189 100644 --- a/tests/fake-codex-fixture.mjs +++ b/tests/fake-codex-fixture.mjs @@ -19,7 +19,7 @@ const readline = require("node:readline"); function loadState() { if (!fs.existsSync(STATE_PATH)) { - return { nextThreadId: 1, nextTurnId: 1, appServerStarts: 0, threads: [], capabilities: null, lastInterrupt: null }; + return { nextThreadId: 1, nextTurnId: 1, appServerStarts: 0, threads: [], subscriptions: [], unsubscribeRequests: [], capabilities: null, lastInterrupt: null }; } return JSON.parse(fs.readFileSync(STATE_PATH, "utf8")); } @@ -313,6 +313,8 @@ rl.on("line", (line) => { throw new Error("thread/start.persistFullHistory requires experimentalApi capability"); } const thread = nextThread(state, message.params.cwd, message.params.ephemeral); + state.subscriptions = [...new Set([...(state.subscriptions || []), thread.id])]; + saveState(state); send({ id: message.id, result: { thread: buildThread(thread), model: message.params.model || "gpt-5.4", modelProvider: "openai", serviceTier: null, cwd: thread.cwd, approvalPolicy: "never", sandbox: { type: "readOnly", access: { type: "fullAccess" }, networkAccess: false }, reasoningEffort: null } }); send({ method: "thread/started", params: { thread: { id: thread.id } } }); break; @@ -346,11 +348,47 @@ rl.on("line", (line) => { } const thread = ensureThread(state, message.params.threadId); thread.updatedAt = now(); + state.subscriptions = [...new Set([...(state.subscriptions || []), thread.id])]; saveState(state); send({ id: message.id, result: { thread: buildThread(thread), model: message.params.model || "gpt-5.4", modelProvider: "openai", serviceTier: null, cwd: thread.cwd, approvalPolicy: "never", sandbox: { type: "readOnly", access: { type: "fullAccess" }, networkAccess: false }, reasoningEffort: null } }); break; } + case "thread/fork": { + const sourceThread = ensureThread(state, message.params.threadId); + const thread = nextThread(state, sourceThread.cwd, message.params.ephemeral); + state.subscriptions = [...new Set([...(state.subscriptions || []), thread.id])]; + saveState(state); + send({ id: message.id, result: { thread: buildThread(thread) } }); + send({ method: "thread/started", params: { thread: { id: thread.id } } }); + break; + } + + case "thread/unsubscribe": { + const subscriptions = state.subscriptions || []; + const wasSubscribed = subscriptions.includes(message.params.threadId); + const wasLoaded = state.threads.some((thread) => thread.id === message.params.threadId); + state.unsubscribeRequests = [...(state.unsubscribeRequests || []), message.params.threadId]; + if (BEHAVIOR === "unsubscribe-fails") { + saveState(state); + send({ id: message.id, error: { code: -32000, message: "thread unsubscribe failed" } }); + break; + } + if (BEHAVIOR === "unsubscribe-notifies") { + send({ + method: "thread/status/changed", + params: { threadId: message.params.threadId, status: { type: "idle" } } + }); + } + state.subscriptions = subscriptions.filter((threadId) => threadId !== message.params.threadId); + saveState(state); + send({ + id: message.id, + result: { status: wasSubscribed ? "unsubscribed" : wasLoaded ? "notSubscribed" : "notLoaded" } + }); + break; + } + case "externalAgentConfig/import": { if (BEHAVIOR === "external-import-unsupported") { send({ id: message.id, error: { code: -32601, message: "Unsupported method: externalAgentConfig/import" } }); @@ -409,6 +447,8 @@ rl.on("line", (line) => { let reviewThread = thread; if (message.params.delivery === "detached") { reviewThread = nextThread(state, thread.cwd, true); + state.subscriptions = [...new Set([...(state.subscriptions || []), reviewThread.id])]; + saveState(state); send({ method: "thread/started", params: { thread: { id: reviewThread.id } } }); } const turnId = nextTurnId(state); @@ -458,6 +498,22 @@ rl.on("line", (line) => { ? structuredReviewPayload(prompt) : taskPayload(prompt, thread.name && thread.name.startsWith("Codex Companion Task") && prompt.includes("Continue from the current thread state")); + if (BEHAVIOR === "with-delayed-subagent") { + setTimeout(() => { + const delayedState = loadState(); + const subThread = nextThread(delayedState, thread.cwd, true); + const subThreadRecord = ensureThread(delayedState, subThread.id); + subThreadRecord.name = "delayed-design-challenger"; + delayedState.subscriptions = [...new Set([...(delayedState.subscriptions || []), subThread.id])]; + saveState(delayedState); + const subTurnId = nextTurnId(delayedState); + send({ method: "thread/started", params: { thread: { ...buildThread(subThreadRecord), name: subThreadRecord.name, agentNickname: subThreadRecord.name } } }); + send({ method: "turn/started", params: { threadId: subThread.id, turn: buildTurn(subTurnId) } }); + send({ method: "turn/completed", params: { threadId: subThread.id, turn: buildTurn(subTurnId, "completed") } }); + }, 100); + break; + } + if ( BEHAVIOR === "with-subagent" || BEHAVIOR === "with-late-subagent-message" || @@ -466,6 +522,7 @@ rl.on("line", (line) => { const subThread = nextThread(state, thread.cwd, true); const subThreadRecord = ensureThread(state, subThread.id); subThreadRecord.name = "design-challenger"; + state.subscriptions = [...new Set([...(state.subscriptions || []), subThread.id])]; saveState(state); const subTurnId = nextTurnId(state); From 1203f5d5228f5fa274163e54c433ce628a710377 Mon Sep 17 00:00:00 2001 From: thossullivan Date: Tue, 1 Sep 2026 13:00:28 -0500 Subject: [PATCH 2/2] fix(broker): harden thread subscription ownership --- plugins/codex/scripts/app-server-broker.mjs | 355 +++++++++++--------- tests/broker-subscriptions.test.mjs | 173 +++++++++- tests/fake-codex-fixture.mjs | 32 +- 3 files changed, 371 insertions(+), 189 deletions(-) diff --git a/plugins/codex/scripts/app-server-broker.mjs b/plugins/codex/scripts/app-server-broker.mjs index 0f420e91f..dbcb068d4 100644 --- a/plugins/codex/scripts/app-server-broker.mjs +++ b/plugins/codex/scripts/app-server-broker.mjs @@ -11,6 +11,7 @@ import { parseBrokerEndpoint } from "./lib/broker-endpoint.mjs"; const STREAMING_METHODS = new Set(["turn/start", "review/start", "thread/compact/start"]); const SUBSCRIBING_METHODS = new Set(["thread/start", "thread/resume", "thread/fork"]); +const UNSUBSCRIBE_RETRY_DELAYS_MS = [100, 500, 2000]; function buildSubscriptionThreadIds(method, result) { const threadIds = new Set(); @@ -34,23 +35,22 @@ function buildProvisionalSubscriptionThreadIds(method, params) { return threadIds; } -function buildNotificationSubscriptionThreadIds(message) { - const threadIds = new Set(); - const params = message?.params; - if (message?.method === "thread/started" && params?.thread?.id) { - threadIds.add(params.thread.id); - } - if (params?.threadId) { - threadIds.add(params.threadId); +function buildNotificationSubscriptionRelationships(message) { + const relationships = []; + const thread = message?.method === "thread/started" ? message.params?.thread : null; + if (thread?.id && thread?.parentThreadId) { + relationships.push({ sourceThreadId: thread.parentThreadId, subscribedThreadId: thread.id }); } - if (Array.isArray(params?.item?.receiverThreadIds)) { - for (const threadId of params.item.receiverThreadIds) { + + const item = message?.params?.item; + if (item?.type === "collabAgentToolCall" && item?.senderThreadId && Array.isArray(item.receiverThreadIds)) { + for (const threadId of item.receiverThreadIds) { if (threadId) { - threadIds.add(threadId); + relationships.push({ sourceThreadId: item.senderThreadId, subscribedThreadId: threadId }); } } } - return threadIds; + return relationships; } function buildStreamThreadIds(method, params, result) { @@ -117,9 +117,19 @@ async function main() { const socketThreadIds = new Map(); const threadSockets = new Map(); const pendingUnsubscribes = new Map(); - const socketUnsubscribingThreadIds = new Map(); + const unsubscribeRetryTimers = new Map(); + + function cancelUnsubscribeRetry(threadId) { + const retry = unsubscribeRetryTimers.get(threadId); + if (!retry) { + return; + } + clearTimeout(retry.timer); + unsubscribeRetryTimers.delete(threadId); + } function addThreadOwner(socket, threadId) { + cancelUnsubscribeRetry(threadId); let ownedThreadIds = socketThreadIds.get(socket); if (!ownedThreadIds) { ownedThreadIds = new Set(); @@ -156,22 +166,6 @@ async function main() { return true; } - function setSocketThreadUnsubscribing(socket, threadId, isUnsubscribing) { - let threadIds = socketUnsubscribingThreadIds.get(socket); - if (isUnsubscribing) { - if (!threadIds) { - threadIds = new Set(); - socketUnsubscribingThreadIds.set(socket, threadIds); - } - threadIds.add(threadId); - return; - } - threadIds?.delete(threadId); - if (threadIds?.size === 0) { - socketUnsubscribingThreadIds.delete(socket); - } - } - function requestThreadUnsubscribe(threadId) { const pending = pendingUnsubscribes.get(threadId); if (pending) { @@ -195,11 +189,35 @@ async function main() { return request; } - async function unsubscribeIfUnowned(threadId) { + function scheduleUnsubscribeRetry(threadId, retryIndex) { + if ( + retryIndex >= UNSUBSCRIBE_RETRY_DELAYS_MS.length || + unsubscribeRetryTimers.has(threadId) || + threadSockets.has(threadId) || + appClient.closed + ) { + return; + } + const timer = setTimeout(() => { + unsubscribeRetryTimers.delete(threadId); + void unsubscribeIfUnowned(threadId, { retryOnFailure: true, retryIndex: retryIndex + 1 }); + }, UNSUBSCRIBE_RETRY_DELAYS_MS[retryIndex]); + timer.unref?.(); + unsubscribeRetryTimers.set(threadId, { timer, retryIndex }); + } + + async function unsubscribeIfUnowned(threadId, { retryOnFailure = false, retryIndex = 0 } = {}) { if (threadSockets.has(threadId) || appClient.closed) { + cancelUnsubscribeRetry(threadId); return null; } - return requestThreadUnsubscribe(threadId); + const outcome = await requestThreadUnsubscribe(threadId); + if (outcome.error && retryOnFailure) { + scheduleUnsubscribeRetry(threadId, retryIndex); + } else if (!outcome.error) { + cancelUnsubscribeRetry(threadId); + } + return outcome; } async function releaseThreadOwners(socket, threadIds = socketThreadIds.get(socket) ?? new Set()) { @@ -209,32 +227,35 @@ async function main() { releasedThreadIds.push(threadId); } } - await Promise.all(releasedThreadIds.map((threadId) => unsubscribeIfUnowned(threadId))); + await Promise.all( + releasedThreadIds.map((threadId) => unsubscribeIfUnowned(threadId, { retryOnFailure: true })) + ); } - function trackSubscriptionResults(socket, method, result, provisionalThreadIds) { + function trackSubscriptionResults(socket, method, result) { for (const threadId of buildSubscriptionThreadIds(method, result)) { - if (provisionalThreadIds.has(threadId)) { - continue; - } if (socket.destroyed || !sockets.has(socket)) { - void unsubscribeIfUnowned(threadId); + void unsubscribeIfUnowned(threadId, { retryOnFailure: true }); continue; } addThreadOwner(socket, threadId); } } - function trackNotificationSubscriptions(socket, message) { - // App-server auto-subscribes its connection to child threads created by subagents. - // Attribute those notification-only subscriptions to the active downstream client. - for (const threadId of buildNotificationSubscriptionThreadIds(message)) { - if (socket && !socket.destroyed && sockets.has(socket)) { - if (!socketUnsubscribingThreadIds.get(socket)?.has(threadId)) { - addThreadOwner(socket, threadId); - } - } else { - void unsubscribeIfUnowned(threadId); + function trackNotificationSubscriptions(message) { + // App-server can auto-subscribe its connection to subagent threads. Attribute + // each child to the downstream owners of its causal parent, not the client + // that happens to be active when a delayed notification arrives. + for (const { sourceThreadId, subscribedThreadId } of buildNotificationSubscriptionRelationships(message)) { + const sourceOwners = [...(threadSockets.get(sourceThreadId) ?? [])].filter( + (socket) => !socket.destroyed && sockets.has(socket) + ); + if (sourceOwners.length === 0) { + void unsubscribeIfUnowned(subscribedThreadId, { retryOnFailure: true }); + continue; + } + for (const socket of sourceOwners) { + addThreadOwner(socket, subscribedThreadId); } } } @@ -250,16 +271,11 @@ async function main() { if (threadSockets.has(threadId)) { return { status: "notSubscribed" }; } - setSocketThreadUnsubscribing(socket, threadId, true); - try { - const outcome = await requestThreadUnsubscribe(threadId); - if (outcome.error) { - throw outcome.error; - } - return outcome.result; - } finally { - setSocketThreadUnsubscribing(socket, threadId, false); + const outcome = await unsubscribeIfUnowned(threadId); + if (outcome?.error) { + throw outcome.error; } + return outcome?.result ?? { status: "notSubscribed" }; } removeThreadOwner(socket, threadId); @@ -267,19 +283,14 @@ async function main() { return { status: "unsubscribed" }; } - setSocketThreadUnsubscribing(socket, threadId, true); - try { - const outcome = await unsubscribeIfUnowned(threadId); - if (outcome?.error) { - if (!socket.destroyed && sockets.has(socket)) { - addThreadOwner(socket, threadId); - } - throw outcome.error; + const outcome = await unsubscribeIfUnowned(threadId); + if (outcome?.error) { + if (!socket.destroyed && sockets.has(socket)) { + addThreadOwner(socket, threadId); } - return outcome?.result ?? { status: "unsubscribed" }; - } finally { - setSocketThreadUnsubscribing(socket, threadId, false); + throw outcome.error; } + return outcome?.result ?? { status: "unsubscribed" }; } function clearSocketOwnership(socket) { @@ -294,7 +305,7 @@ async function main() { function routeNotification(message) { const target = activeRequestSocket ?? activeStreamSocket; - trackNotificationSubscriptions(target, message); + trackNotificationSubscriptions(message); if (!target) { return; } @@ -312,6 +323,10 @@ async function main() { } async function shutdown(server) { + for (const { timer } of unsubscribeRetryTimers.values()) { + clearTimeout(timer); + } + unsubscribeRetryTimers.clear(); for (const socket of sockets) { socket.end(); } @@ -331,134 +346,140 @@ async function main() { sockets.add(socket); socket.setEncoding("utf8"); let buffer = ""; + let processing = Promise.resolve(); - socket.on("data", async (chunk) => { - buffer += chunk; - let newlineIndex = buffer.indexOf("\n"); - while (newlineIndex !== -1) { - const line = buffer.slice(0, newlineIndex); - buffer = buffer.slice(newlineIndex + 1); - newlineIndex = buffer.indexOf("\n"); - - if (!line.trim()) { - continue; - } + async function handleLine(line) { + if (!line.trim() || socket.destroyed || !sockets.has(socket)) { + return; + } - let message; - try { - message = JSON.parse(line); - } catch (error) { - send(socket, { - id: null, - error: buildJsonRpcError(-32700, `Invalid JSON: ${error.message}`) - }); - continue; - } + let message; + try { + message = JSON.parse(line); + } catch (error) { + send(socket, { + id: null, + error: buildJsonRpcError(-32700, `Invalid JSON: ${error.message}`) + }); + return; + } - if (message.id !== undefined && message.method === "initialize") { - send(socket, { - id: message.id, - result: { - userAgent: "codex-companion-broker" - } - }); - continue; - } + if (message.id !== undefined && message.method === "initialize") { + send(socket, { + id: message.id, + result: { + userAgent: "codex-companion-broker" + } + }); + return; + } - if (message.method === "initialized" && message.id === undefined) { - continue; - } + if (message.method === "initialized" && message.id === undefined) { + return; + } - if (message.id !== undefined && message.method === "broker/shutdown") { - send(socket, { id: message.id, result: {} }); - await shutdown(server); - process.exit(0); - } + if (message.id !== undefined && message.method === "broker/shutdown") { + send(socket, { id: message.id, result: {} }); + await shutdown(server); + process.exit(0); + } - if (message.id === undefined) { - continue; - } + if (message.id === undefined) { + return; + } - const allowInterruptDuringActiveStream = - isInterruptRequest(message) && activeStreamSocket && activeStreamSocket !== socket && !activeRequestSocket; + const allowInterruptDuringActiveStream = + isInterruptRequest(message) && activeStreamSocket && activeStreamSocket !== socket && !activeRequestSocket; + + if ( + ((activeRequestSocket && activeRequestSocket !== socket) || (activeStreamSocket && activeStreamSocket !== socket)) && + !allowInterruptDuringActiveStream + ) { + send(socket, { + id: message.id, + error: buildJsonRpcError(BROKER_BUSY_RPC_CODE, "Shared Codex broker is busy.") + }); + return; + } - if ( - ((activeRequestSocket && activeRequestSocket !== socket) || (activeStreamSocket && activeStreamSocket !== socket)) && - !allowInterruptDuringActiveStream - ) { + if (allowInterruptDuringActiveStream) { + try { + const result = await appClient.request(message.method, message.params ?? {}); + send(socket, { id: message.id, result }); + } catch (error) { send(socket, { id: message.id, - error: buildJsonRpcError(BROKER_BUSY_RPC_CODE, "Shared Codex broker is busy.") + error: buildJsonRpcError(error.rpcCode ?? -32000, error.message) }); - continue; } + return; + } - if (allowInterruptDuringActiveStream) { - try { - const result = await appClient.request(message.method, message.params ?? {}); - send(socket, { id: message.id, result }); - } catch (error) { - send(socket, { - id: message.id, - error: buildJsonRpcError(error.rpcCode ?? -32000, error.message) - }); - } - continue; + const isStreaming = STREAMING_METHODS.has(message.method); + // Claim known thread ids before awaiting app-server. This prevents another + // client's close handler from unsubscribing a concurrently resumed thread. + const provisionalThreadIds = buildProvisionalSubscriptionThreadIds(message.method, message.params ?? {}); + const addedProvisionalThreadIds = new Set(); + for (const threadId of provisionalThreadIds) { + if (addThreadOwner(socket, threadId)) { + addedProvisionalThreadIds.add(threadId); } + } + activeRequestSocket = socket; - const isStreaming = STREAMING_METHODS.has(message.method); - // Claim known thread ids before awaiting app-server. This prevents another - // client's close handler from unsubscribing a concurrently resumed thread. - const provisionalThreadIds = buildProvisionalSubscriptionThreadIds(message.method, message.params ?? {}); - const addedProvisionalThreadIds = new Set(); - for (const threadId of provisionalThreadIds) { - if (addThreadOwner(socket, threadId)) { - addedProvisionalThreadIds.add(threadId); - } + try { + const result = + message.method === "thread/unsubscribe" + ? await handleThreadUnsubscribe(socket, message.params ?? {}) + : await appClient.request(message.method, message.params ?? {}); + trackSubscriptionResults(socket, message.method, result); + send(socket, { id: message.id, result }); + if (isStreaming && !socket.destroyed && sockets.has(socket)) { + activeStreamSocket = socket; + activeStreamThreadIds = buildStreamThreadIds(message.method, message.params ?? {}, result); } - activeRequestSocket = socket; - - try { - const result = - message.method === "thread/unsubscribe" - ? await handleThreadUnsubscribe(socket, message.params ?? {}) - : await appClient.request(message.method, message.params ?? {}); - trackSubscriptionResults(socket, message.method, result, provisionalThreadIds); - send(socket, { id: message.id, result }); - if (isStreaming) { - activeStreamSocket = socket; - activeStreamThreadIds = buildStreamThreadIds(message.method, message.params ?? {}, result); - } - if (activeRequestSocket === socket) { - activeRequestSocket = null; - } - } catch (error) { - await releaseThreadOwners(socket, addedProvisionalThreadIds); - send(socket, { - id: message.id, - error: buildJsonRpcError(error.rpcCode ?? -32000, error.message) - }); - if (activeRequestSocket === socket) { - activeRequestSocket = null; - } - if (activeStreamSocket === socket && !isStreaming) { - activeStreamSocket = null; - } + if (activeRequestSocket === socket) { + activeRequestSocket = null; } + } catch (error) { + await releaseThreadOwners(socket, addedProvisionalThreadIds); + send(socket, { + id: message.id, + error: buildJsonRpcError(error.rpcCode ?? -32000, error.message) + }); + if (activeRequestSocket === socket) { + activeRequestSocket = null; + } + if (activeStreamSocket === socket && !isStreaming) { + activeStreamSocket = null; + } + } + } + + socket.on("data", (chunk) => { + buffer += chunk; + let newlineIndex = buffer.indexOf("\n"); + while (newlineIndex !== -1) { + const line = buffer.slice(0, newlineIndex); + buffer = buffer.slice(newlineIndex + 1); + newlineIndex = buffer.indexOf("\n"); + processing = processing.then(() => handleLine(line)).catch((error) => { + process.stderr.write( + `Failed to process broker request: ${error instanceof Error ? error.message : String(error)}\n` + ); + }); } }); socket.on("close", () => { sockets.delete(socket); clearSocketOwnership(socket); - socketUnsubscribingThreadIds.delete(socket); void releaseThreadOwners(socket); }); socket.on("error", () => { sockets.delete(socket); clearSocketOwnership(socket); - socketUnsubscribingThreadIds.delete(socket); void releaseThreadOwners(socket); }); }); diff --git a/tests/broker-subscriptions.test.mjs b/tests/broker-subscriptions.test.mjs index 3277b611d..88fa48a9e 100644 --- a/tests/broker-subscriptions.test.mjs +++ b/tests/broker-subscriptions.test.mjs @@ -16,6 +16,20 @@ function delay(ms) { return new Promise((resolve) => setTimeout(resolve, ms)); } +async function waitWithTimeout(promise, timeoutMs) { + let timer; + try { + return await Promise.race([ + promise, + new Promise((resolve) => { + timer = setTimeout(() => resolve(null), timeoutMs); + }) + ]); + } finally { + clearTimeout(timer); + } +} + async function waitFor(predicate, { timeoutMs = 10000, intervalMs = 25 } = {}) { const startedAt = Date.now(); while (Date.now() - startedAt < timeoutMs) { @@ -43,6 +57,7 @@ function startBroker(behavior = "review-ok") { const socketPath = path.join(sessionDir, "broker.sock"); const pidFile = path.join(sessionDir, "broker.pid"); const statePath = path.join(binDir, "fake-codex-state.json"); + const tempDirs = [binDir, sessionDir, cwd]; const child = spawn( process.execPath, @@ -58,15 +73,21 @@ function startBroker(behavior = "review-ok") { }); async function stop() { - if (child.exitCode !== null || child.signalCode !== null) { - await exited; - return; - } - child.kill("SIGTERM"); - const result = await Promise.race([exited, delay(5000).then(() => null)]); - if (!result) { - child.kill("SIGKILL"); - await exited; + try { + if (child.exitCode !== null || child.signalCode !== null) { + await exited; + return; + } + child.kill("SIGTERM"); + const result = await waitWithTimeout(exited, 5000); + if (!result) { + child.kill("SIGKILL"); + await exited; + } + } finally { + for (const tempDir of tempDirs) { + fs.rmSync(tempDir, { recursive: true, force: true }); + } } } @@ -97,8 +118,7 @@ async function connectClient(socketPath) { function notifyWaiters(message) { for (const waiter of [...notificationWaiters]) { if (waiter.predicate(message)) { - notificationWaiters.delete(waiter); - waiter.resolve(message); + waiter.settle(message); } } } @@ -152,10 +172,19 @@ async function connectClient(socketPath) { if (existing) { return existing; } - const notification = new Promise((resolve) => { - notificationWaiters.add({ predicate, resolve }); + return new Promise((resolve) => { + let timer; + const waiter = { + predicate, + settle(message) { + clearTimeout(timer); + notificationWaiters.delete(waiter); + resolve(message); + } + }; + timer = setTimeout(() => waiter.settle(null), timeoutMs); + notificationWaiters.add(waiter); }); - return Promise.race([notification, delay(timeoutMs).then(() => null)]); } await request("initialize", {}); @@ -296,6 +325,30 @@ test("broker unsubscribes auto-subscribed subagent threads", async (t) => { assert.deepEqual(readState(broker.statePath).subscriptions, []); }); +test("broker tracks subagent subscriptions from collaboration items", async (t) => { + const broker = startBroker("with-receiver-only-subagent"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const threadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: true })).thread.id; + await client.request("turn/start", { + threadId, + input: [{ type: "text", text: "delegate without a thread-started notification" }] + }); + const completed = await client.waitForNotification( + (message) => message.method === "turn/completed" && message.params?.threadId === threadId + ); + assert.ok(completed, "task never completed"); + + const subagentThread = readState(broker.statePath).threads.find((thread) => thread.name === "design-challenger"); + assert.ok(subagentThread, "subagent thread was not created"); + + await client.end(); + await waitForUnsubscribes(broker.statePath, [threadId, subagentThread.id]); + assert.deepEqual(readState(broker.statePath).subscriptions, []); +}); + test("broker unsubscribes a child thread created after its client disconnects", async (t) => { const broker = startBroker("with-delayed-subagent"); t.after(() => broker.stop()); @@ -321,6 +374,80 @@ test("broker unsubscribes a child thread created after its client disconnects", assert.deepEqual(readState(broker.statePath).subscriptions, []); }); +test("broker does not assign a delayed child thread to an unrelated active client", async (t) => { + const broker = startBroker("with-delayed-subagent"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const firstClient = await connectClient(broker.socketPath); + const firstThreadId = ( + await firstClient.request("thread/start", { cwd: process.cwd(), ephemeral: true }) + ).thread.id; + await firstClient.request("turn/start", { + threadId: firstThreadId, + input: [{ type: "text", text: "create a delayed child" }] + }); + firstClient.destroy(); + + const secondClient = await connectClient(broker.socketPath); + const secondThreadId = ( + await secondClient.request("thread/start", { cwd: process.cwd(), ephemeral: true }) + ).thread.id; + await secondClient.request("turn/start", { + threadId: secondThreadId, + input: [{ type: "text", text: "remain active while the first child arrives" }] + }); + + const childrenCreated = await waitFor( + () => readState(broker.statePath)?.threads.filter((thread) => thread.parentThreadId).length === 2 + ); + assert.equal(childrenCreated, true, "delayed child threads were not created"); + + const state = readState(broker.statePath); + const firstChild = state.threads.find((thread) => thread.parentThreadId === firstThreadId); + const secondChild = state.threads.find((thread) => thread.parentThreadId === secondThreadId); + assert.ok(firstChild, "the first client's child thread was not recorded"); + assert.ok(secondChild, "the second client's child thread was not recorded"); + + const firstReleased = await waitFor(() => { + const requests = readState(broker.statePath)?.unsubscribeRequests ?? []; + return requests.includes(firstThreadId) && requests.includes(firstChild.id); + }); + assert.equal(firstReleased, true, "the disconnected client's subscriptions were not released"); + const requestsBeforeSecondClientCloses = readState(broker.statePath).unsubscribeRequests; + assert.equal(requestsBeforeSecondClientCloses.includes(secondThreadId), false); + assert.equal(requestsBeforeSecondClientCloses.includes(secondChild.id), false); + + await secondClient.end(); + await waitForUnsubscribes(broker.statePath, [firstThreadId, firstChild.id, secondThreadId, secondChild.id]); + assert.deepEqual(readState(broker.statePath).subscriptions, []); +}); + +test("broker serializes subscription requests from one downstream client", async (t) => { + const broker = startBroker("overlapping-resume"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const firstClient = await connectClient(broker.socketPath); + const threadId = (await firstClient.request("thread/start", { cwd: process.cwd(), ephemeral: false })).thread.id; + const secondClient = await connectClient(broker.socketPath); + const results = await Promise.allSettled([ + secondClient.request("thread/resume", { threadId, persistFullHistory: true }), + secondClient.request("thread/resume", { threadId }) + ]); + assert.equal(results[0].status, "rejected"); + assert.match(results[0].reason.message, /forced delayed resume failure/); + assert.equal(results[1].status, "fulfilled"); + + await firstClient.end(); + await delay(250); + assert.deepEqual(readState(broker.statePath).unsubscribeRequests, []); + assert.deepEqual(readState(broker.statePath).subscriptions, [threadId]); + + await secondClient.end(); + await waitForUnsubscribes(broker.statePath, [threadId]); +}); + test("broker keeps shared upstream subscriptions when one client explicitly unsubscribes", async (t) => { const broker = startBroker(); t.after(() => broker.stop()); @@ -378,3 +505,21 @@ test("broker logs upstream unsubscribe failures", async (t) => { assert.equal(warningObserved, true, "unsubscribe failure was not logged"); assert.deepEqual(readState(broker.statePath).subscriptions, [threadId]); }); + +test("broker retries a transient upstream unsubscribe failure", async (t) => { + const broker = startBroker("unsubscribe-fails-once"); + t.after(() => broker.stop()); + assert.equal(await broker.listening(), true, `broker never listened: ${broker.stderr()}`); + + const client = await connectClient(broker.socketPath); + const threadId = (await client.request("thread/start", { cwd: process.cwd(), ephemeral: true })).thread.id; + await client.end(); + + const retried = await waitFor(() => { + const state = readState(broker.statePath); + return state?.unsubscribeRequests?.length === 2 && state.subscriptions.length === 0; + }); + assert.equal(retried, true, "the failed unsubscribe was not retried"); + assert.deepEqual(readState(broker.statePath).unsubscribeRequests, [threadId, threadId]); + assert.match(broker.stderr(), new RegExp(`Failed to unsubscribe Codex thread ${threadId}`)); +}); diff --git a/tests/fake-codex-fixture.mjs b/tests/fake-codex-fixture.mjs index db0efc189..125f1b897 100644 --- a/tests/fake-codex-fixture.mjs +++ b/tests/fake-codex-fixture.mjs @@ -42,6 +42,8 @@ function now() { function buildThread(thread) { return { id: thread.id, + forkedFromId: thread.forkedFromId || null, + parentThreadId: thread.parentThreadId || null, preview: thread.preview || "", ephemeral: Boolean(thread.ephemeral), modelProvider: "openai", @@ -116,9 +118,11 @@ function send(message) { process.stdout.write(JSON.stringify(message) + "\\n"); } -function nextThread(state, cwd, ephemeral) { +function nextThread(state, cwd, ephemeral, { forkedFromId = null, parentThreadId = null } = {}) { const thread = { id: "thr_" + state.nextThreadId++, + forkedFromId, + parentThreadId, cwd: cwd || process.cwd(), name: null, preview: "", @@ -316,7 +320,7 @@ rl.on("line", (line) => { state.subscriptions = [...new Set([...(state.subscriptions || []), thread.id])]; saveState(state); send({ id: message.id, result: { thread: buildThread(thread), model: message.params.model || "gpt-5.4", modelProvider: "openai", serviceTier: null, cwd: thread.cwd, approvalPolicy: "never", sandbox: { type: "readOnly", access: { type: "fullAccess" }, networkAccess: false }, reasoningEffort: null } }); - send({ method: "thread/started", params: { thread: { id: thread.id } } }); + send({ method: "thread/started", params: { thread: buildThread(thread) } }); break; } @@ -343,6 +347,12 @@ rl.on("line", (line) => { } case "thread/resume": { + if (BEHAVIOR === "overlapping-resume" && message.params.persistFullHistory === true) { + setTimeout(() => { + send({ id: message.id, error: { code: -32000, message: "forced delayed resume failure" } }); + }, 100); + break; + } if (requiresExperimental("persistExtendedHistory", message, state) || requiresExperimental("persistFullHistory", message, state)) { throw new Error("thread/resume.persistFullHistory requires experimentalApi capability"); } @@ -356,11 +366,11 @@ rl.on("line", (line) => { case "thread/fork": { const sourceThread = ensureThread(state, message.params.threadId); - const thread = nextThread(state, sourceThread.cwd, message.params.ephemeral); + const thread = nextThread(state, sourceThread.cwd, message.params.ephemeral, { forkedFromId: sourceThread.id }); state.subscriptions = [...new Set([...(state.subscriptions || []), thread.id])]; saveState(state); send({ id: message.id, result: { thread: buildThread(thread) } }); - send({ method: "thread/started", params: { thread: { id: thread.id } } }); + send({ method: "thread/started", params: { thread: buildThread(thread) } }); break; } @@ -369,7 +379,10 @@ rl.on("line", (line) => { const wasSubscribed = subscriptions.includes(message.params.threadId); const wasLoaded = state.threads.some((thread) => thread.id === message.params.threadId); state.unsubscribeRequests = [...(state.unsubscribeRequests || []), message.params.threadId]; - if (BEHAVIOR === "unsubscribe-fails") { + if ( + BEHAVIOR === "unsubscribe-fails" || + (BEHAVIOR === "unsubscribe-fails-once" && state.unsubscribeRequests.length === 1) + ) { saveState(state); send({ id: message.id, error: { code: -32000, message: "thread unsubscribe failed" } }); break; @@ -501,7 +514,7 @@ rl.on("line", (line) => { if (BEHAVIOR === "with-delayed-subagent") { setTimeout(() => { const delayedState = loadState(); - const subThread = nextThread(delayedState, thread.cwd, true); + const subThread = nextThread(delayedState, thread.cwd, true, { parentThreadId: thread.id }); const subThreadRecord = ensureThread(delayedState, subThread.id); subThreadRecord.name = "delayed-design-challenger"; delayedState.subscriptions = [...new Set([...(delayedState.subscriptions || []), subThread.id])]; @@ -516,17 +529,20 @@ rl.on("line", (line) => { if ( BEHAVIOR === "with-subagent" || + BEHAVIOR === "with-receiver-only-subagent" || BEHAVIOR === "with-late-subagent-message" || BEHAVIOR === "with-subagent-no-main-turn-completed" ) { - const subThread = nextThread(state, thread.cwd, true); + const subThread = nextThread(state, thread.cwd, true, { parentThreadId: thread.id }); const subThreadRecord = ensureThread(state, subThread.id); subThreadRecord.name = "design-challenger"; state.subscriptions = [...new Set([...(state.subscriptions || []), subThread.id])]; saveState(state); const subTurnId = nextTurnId(state); - send({ method: "thread/started", params: { thread: { ...buildThread(subThreadRecord), name: "design-challenger", agentNickname: "design-challenger" } } }); + if (BEHAVIOR !== "with-receiver-only-subagent") { + send({ method: "thread/started", params: { thread: { ...buildThread(subThreadRecord), name: "design-challenger", agentNickname: "design-challenger" } } }); + } send({ method: "turn/started", params: { threadId: thread.id, turn: buildTurn(turnId) } }); send({ method: "item/started",