fix(computer-use): keep startup enumeration out of turns

This commit is contained in:
程序员阿江(Relakkes)
2026-09-02 10:42:42 +08:00
parent 3e3ae6510a
commit 19e88db8ea
7 changed files with 310 additions and 6 deletions
@@ -28,7 +28,7 @@ enum HelperClientPolicy {
private static let daemonCommands: Set<String> = [
// Codex semantic tool contract. `resolve_app_target` is the mandatory
// read-only authorization preflight used before the ten public tools.
"list_apps", "resolve_app_target", "get_app_state", "click",
"list_apps", "list_installed_apps", "resolve_app_target", "get_app_state", "click",
"set_value", "select_text", "perform_secondary_action", "scroll",
"drag", "press_key", "type_text", "paste",
// Private lifecycle / visible-overlay / permission and input diagnostics.
@@ -10,7 +10,7 @@ enum ComputerUseDaemonProtocol {
static let version = "CCHahaComputerUseIPC-2"
static let maxFrameBytes = 8 * 1024 * 1024
private static let connectionScopedCommands: Set<String> = [
"ping", "check_permissions", "shutdown",
"ping", "check_permissions", "list_installed_apps", "shutdown",
]
static func isTurnScoped(command: String) -> Bool {
@@ -62,7 +62,7 @@ enum ComputerUseDaemonProtocol {
}
/// Enforces one explicit turn at a time on the single authenticated connection.
/// `ping` negotiates the protocol without opening a turn. Every other request
/// Connection setup and diagnostics do not open a turn. Every semantic request
/// must keep the same identity until `turn_end` releases all turn-owned state.
struct DaemonTurnGate {
private(set) var active: ComputerUseDaemonProtocol.Metadata?
@@ -363,7 +363,7 @@ struct ClientAttestationTests {
@Test("daemon exposes only the semantic contract and required control diagnostics")
func daemonCommandAllowlist() {
let allowed = [
"list_apps", "resolve_app_target", "get_app_state", "click",
"list_apps", "list_installed_apps", "resolve_app_target", "get_app_state", "click",
"set_value", "select_text", "perform_secondary_action", "scroll",
"drag", "press_key", "type_text", "paste", "ping", "shutdown",
"overlay_show", "overlay_hide", "turn_end", "check_permissions",
@@ -377,7 +377,7 @@ struct ClientAttestationTests {
"screenshot", "resolve_prepare_capture", "zoom", "key", "type",
"hold_key", "paste_clipboard", "read_clipboard", "write_clipboard",
"move_mouse", "mouse_down", "mouse_up", "cursor_position",
"open_app", "list_installed_apps", "list_running_apps",
"open_app", "list_running_apps",
] {
#expect(!HelperClientPolicy.isDaemonCommandAllowed(command))
}
@@ -113,4 +113,22 @@ final class DaemonProtocolTests: XCTestCase {
XCTAssertNil(gate.active)
}
func testStartupAppEnumerationDoesNotPoisonResumedSessionTurn() throws {
let bootstrap = ComputerUseDaemonProtocol.Metadata(
sessionId: "bootstrap-session",
turnId: "connection-1"
)
let resumed = ComputerUseDaemonProtocol.Metadata(
sessionId: "resumed-session",
turnId: "turn-a"
)
var gate = DaemonTurnGate()
try gate.admit(bootstrap, command: "list_installed_apps")
XCTAssertNil(gate.active)
try gate.admit(resumed, command: "get_app_state")
XCTAssertEqual(gate.active, resumed)
}
}
@@ -0,0 +1,158 @@
import { afterEach, describe, expect, test } from 'bun:test'
import { spawnSync } from 'node:child_process'
import fs from 'node:fs'
import net from 'node:net'
import os from 'node:os'
import path from 'node:path'
import { fileURLToPath } from 'node:url'
import { getSessionId, switchSession } from '../../bootstrap/state.js'
import { asSessionId } from '../../types/ids.js'
import { runCleanupFunctions } from '../cleanupRegistry.js'
import {
__resetDaemonClientForTests,
__setDaemonSocketForTests,
callDaemon,
} from './cuHelperDaemon.js'
type EchoedRequest = {
cmd: string
sessionId: string
turnId: string
}
const fixture = path.join(
path.dirname(fileURLToPath(import.meta.url)),
'test-fixtures',
'launchdDaemon.mjs',
)
const spawnedPids = new Set<number>()
const sockets = new Set<net.Socket>()
let runtimeRoot: string | undefined
const originalSessionId = getSessionId()
function isAlive(pid: number): boolean {
try {
process.kill(pid, 0)
return true
} catch {
return false
}
}
async function waitUntil(
predicate: () => boolean,
description: string,
timeoutMs = 3_000,
): Promise<void> {
const deadline = Date.now() + timeoutMs
while (Date.now() < deadline) {
if (predicate()) return
await Bun.sleep(20)
}
throw new Error(`Timed out waiting for ${description}`)
}
async function launchDetachedDaemon(socketPath: string): Promise<number> {
const result = spawnSync(process.execPath, [fixture, 'launch', socketPath], {
encoding: 'utf8',
timeout: 5_000,
})
if (result.error || result.status !== 0) {
throw result.error ?? new Error(result.stderr || `launcher exited ${result.status}`)
}
const pid = Number.parseInt(result.stdout.trim(), 10)
if (!Number.isSafeInteger(pid) || pid <= 0) {
throw new Error(`invalid detached daemon pid: ${result.stdout}`)
}
spawnedPids.add(pid)
await waitUntil(
() => fs.existsSync(socketPath) && fs.existsSync(`${socketPath}.pid`),
'detached daemon endpoints',
)
const parent = spawnSync('/bin/ps', ['-o', 'ppid=', '-p', String(pid)], {
encoding: 'utf8',
})
expect(parent.status).toBe(0)
expect(Number.parseInt(parent.stdout.trim(), 10)).toBe(1)
return pid
}
async function connect(socketPath: string): Promise<net.Socket> {
const socket = net.createConnection(socketPath)
sockets.add(socket)
await new Promise<void>((resolve, reject) => {
socket.once('connect', resolve)
socket.once('error', reject)
})
return socket
}
async function driveResumeBoundary(
socketPath: string,
bootstrapSession: string,
resumedSession: string,
): Promise<void> {
const daemonPid = await launchDetachedDaemon(socketPath)
const socket = await connect(socketPath)
__setDaemonSocketForTests(socket, 2_000)
switchSession(asSessionId(bootstrapSession))
const enumeration = await callDaemon<EchoedRequest>('list_installed_apps')
expect(enumeration).toMatchObject({
cmd: 'list_installed_apps',
sessionId: bootstrapSession,
})
expect(enumeration.turnId).toMatch(/^connection-/)
switchSession(asSessionId(resumedSession))
const firstState = await callDaemon<EchoedRequest>('get_app_state', {
app: 'TextEdit',
})
expect(firstState).toMatchObject({
cmd: 'get_app_state',
sessionId: resumedSession,
})
expect(firstState.turnId).not.toBe(enumeration.turnId)
await runCleanupFunctions()
await waitUntil(
() => !isAlive(daemonPid)
&& !fs.existsSync(socketPath)
&& !fs.existsSync(`${socketPath}.pid`),
'socket-owned daemon shutdown',
)
spawnedPids.delete(daemonPid)
sockets.delete(socket)
}
afterEach(async () => {
try { await runCleanupFunctions() } catch {}
__resetDaemonClientForTests()
switchSession(originalSessionId)
for (const socket of sockets) socket.destroy()
sockets.clear()
for (const pid of spawnedPids) {
if (isAlive(pid)) {
try { process.kill(pid, 'SIGTERM') } catch {}
}
}
spawnedPids.clear()
if (runtimeRoot) fs.rmSync(runtimeRoot, { recursive: true, force: true })
runtimeRoot = undefined
})
describe.skipIf(process.platform !== 'darwin')(
'cu-helper launchd-owned daemon lifecycle',
() => {
test('startup enumeration cannot poison the first resumed turn across shutdown and restart', async () => {
runtimeRoot = fs.mkdtempSync(path.join(os.tmpdir(), 'cc-haha-cu-process-'))
const socketPath = path.join(runtimeRoot, 'cu-helper.sock')
const resumedSession = 'resumed-session'
await driveResumeBoundary(socketPath, 'bootstrap-before-quit', resumedSession)
await driveResumeBoundary(socketPath, 'bootstrap-after-restart', resumedSession)
})
},
)
+7 -1
View File
@@ -36,6 +36,12 @@ const REQUEST_TIMEOUT_MS = 20_000
const READINESS_TIMEOUT_MS = 8_000
const SHUTDOWN_GRACE_MS = 1_000
export const CU_HELPER_PROTOCOL_VERSION = 'CCHahaComputerUseIPC-2'
const CONNECTION_SCOPED_COMMANDS = new Set([
'ping',
'check_permissions',
'list_installed_apps',
'shutdown',
])
/**
* Thrown ONLY for daemon INFRASTRUCTURE failures known to happen before the
@@ -615,7 +621,7 @@ function dispatchDaemonCommand<T>(
))
}
const id = String(++state.nextId)
const isTurnScoped = !['ping', 'check_permissions', 'shutdown'].includes(command)
const isTurnScoped = !CONNECTION_SCOPED_COMMANDS.has(command)
const turnId = activeTurnId
?? (isTurnScoped ? (activeTurnId = randomUUID()) : `connection-${state.generation}`)
const request = {
@@ -0,0 +1,122 @@
import { spawn } from 'node:child_process'
import fs from 'node:fs'
import net from 'node:net'
import { fileURLToPath } from 'node:url'
const [mode, socketPath] = process.argv.slice(2)
if (!socketPath || (mode !== 'launch' && mode !== 'serve')) {
process.stderr.write('usage: launchdDaemon.mjs <launch|serve> <socket>\n')
process.exit(2)
}
if (mode === 'launch') {
const child = spawn(process.execPath, [fileURLToPath(import.meta.url), 'serve', socketPath], {
detached: true,
stdio: 'ignore',
})
if (!child.pid) {
process.stderr.write('detached daemon did not start\n')
process.exit(1)
}
child.unref()
process.stdout.write(`${child.pid}\n`)
} else {
const pidfile = `${socketPath}.pid`
let activeTurn
let activeSocket
let closing = false
function removeEndpoints() {
try { fs.rmSync(socketPath, { force: true }) } catch {}
try { fs.rmSync(pidfile, { force: true }) } catch {}
}
function terminate() {
if (closing) return
closing = true
activeSocket?.destroy()
server.close(() => {
removeEndpoints()
process.exit(0)
})
setTimeout(() => {
removeEndpoints()
process.exit(0)
}, 1_000).unref()
}
function reply(socket, id, body) {
socket.write(`${JSON.stringify({ id, ...body })}\n`)
}
function handleRequest(socket, request) {
const metadata = {
sessionId: request.sessionId,
turnId: request.turnId,
}
const connectionScoped = request.cmd === 'ping'
|| request.cmd === 'check_permissions'
|| request.cmd === 'shutdown'
|| request.turnId?.startsWith('connection-')
if (!connectionScoped) {
if (
activeTurn
&& (activeTurn.sessionId !== metadata.sessionId || activeTurn.turnId !== metadata.turnId)
) {
reply(socket, request.id, {
ok: false,
error: {
code: 'turn_mismatch',
message: 'A different Computer Use turn is still active; finish it before starting another',
},
})
return
}
activeTurn ??= metadata
}
if (request.cmd === 'turn_end' || request.cmd === 'overlay_hide') {
activeTurn = undefined
}
reply(socket, request.id, {
ok: true,
result: {
cmd: request.cmd,
sessionId: request.sessionId,
turnId: request.turnId,
},
})
if (request.cmd === 'shutdown') {
setImmediate(terminate)
}
}
removeEndpoints()
const server = net.createServer(socket => {
if (activeSocket) {
socket.destroy()
return
}
activeSocket = socket
let buffered = ''
socket.on('data', chunk => {
buffered += chunk.toString()
let newline
while ((newline = buffered.indexOf('\n')) >= 0) {
const line = buffered.slice(0, newline)
buffered = buffered.slice(newline + 1)
if (!line.trim()) continue
handleRequest(socket, JSON.parse(line))
}
})
socket.once('close', terminate)
})
server.listen(socketPath, () => {
fs.chmodSync(socketPath, 0o600)
fs.writeFileSync(pidfile, String(process.pid), { mode: 0o600 })
})
process.once('SIGTERM', terminate)
process.once('SIGINT', terminate)
}