fix(acp): drain updates before end turn (#40422)

This commit is contained in:
Shoubhit Dash
2026-08-04 17:26:33 +05:30
committed by GitHub
parent dc67b7a3be
commit 44614c79c4
3 changed files with 264 additions and 51 deletions
+88 -9
View File
@@ -40,7 +40,10 @@ export class Subscription {
private readonly abort = new AbortController()
private readonly shellSnapshots = new Map<string, string>()
private readonly toolStarts = new Set<string>()
private readonly connectionWaiters = new Set<() => void>()
private readonly idleWaiters = new Map<string, Set<ReturnType<typeof signal>>>()
private readonly permission: ACPPermission.Handler
private connected = false
private started = false
constructor(
@@ -63,10 +66,35 @@ export class Subscription {
stop() {
this.abort.abort()
this.disconnected()
for (const resolve of this.connectionWaiters) resolve()
this.connectionWaiters.clear()
}
async runUntilIdle<A>(sessionId: string, request: () => Promise<A>) {
await this.waitUntilConnected()
const waiter = signal()
const waiters = this.idleWaiters.get(sessionId) ?? new Set()
waiters.add(waiter)
this.idleWaiters.set(sessionId, waiters)
try {
// Idle is queued after the turn's events, and this subscription awaits each update in order.
void waiter.promise.catch(() => {})
const response = await request()
await waiter.promise
return response
} finally {
waiters.delete(waiter)
if (waiters.size === 0) this.idleWaiters.delete(sessionId)
}
}
async handle(event: Event) {
switch (event.type) {
case "session.status":
if (event.properties.status.type === "idle") this.idle(event.properties.sessionID)
return
case "permission.asked":
this.permission.handle(event)
return
@@ -115,19 +143,51 @@ export class Subscription {
private async run() {
while (!this.abort.signal.aborted) {
const events = (await this.input.sdk.global.event({
signal: this.abort.signal,
})) as GlobalEventStream
for await (const event of events.stream) {
if (this.abort.signal.aborted) return
if (!event.payload) continue
await this.handle(event.payload).catch(() => {})
}
await this.consume().catch(() => {})
this.disconnected()
if (!this.abort.signal.aborted) await new Promise((resolve) => setTimeout(resolve, 1000))
}
}
private async consume() {
const events = (await this.input.sdk.global.event({
signal: this.abort.signal,
})) as GlobalEventStream
this.connected = true
for (const resolve of this.connectionWaiters) resolve()
this.connectionWaiters.clear()
for await (const event of events.stream) {
if (this.abort.signal.aborted) return
if (!event.payload) continue
await this.handle(event.payload).catch(() => {})
}
}
private async waitUntilConnected() {
while (!this.connected) {
if (this.abort.signal.aborted) throw new Error("ACP event subscription stopped")
await new Promise<void>((resolve) => this.connectionWaiters.add(resolve))
}
}
private disconnected() {
if (!this.connected) return
this.connected = false
const error = new Error("ACP event stream disconnected")
for (const waiters of this.idleWaiters.values()) {
for (const waiter of waiters) waiter.reject(error)
}
this.idleWaiters.clear()
}
private idle(sessionId: string) {
const waiters = this.idleWaiters.get(sessionId)
if (!waiters) return
this.idleWaiters.delete(sessionId)
for (const waiter of waiters) waiter.resolve()
}
private async handlePartUpdated(event: EventMessagePartUpdated) {
const part = event.properties.part
const sessionId = part.sessionID || event.properties.sessionID
@@ -339,4 +399,23 @@ export class Subscription {
}
}
function signal() {
const state: {
resolve: () => void
reject: (reason?: unknown) => void
} = {
resolve: () => {},
reject: () => {},
}
const promise = new Promise<void>((resolve, reject) => {
state.resolve = resolve
state.reject = reject
})
return {
promise,
resolve: () => state.resolve(),
reject: (reason?: unknown) => state.reject(reason),
}
}
export * as ACPEvent from "./event"
+39 -31
View File
@@ -88,6 +88,8 @@ export function make(input: {
? ACPEvent.start({ sdk: input.sdk, connection: input.connection, session })
: undefined
if (events) input.eventSubscription?.(events)
const runUntilIdle = <A>(sessionId: string, fn: () => Promise<A>) =>
events ? events.runUntilIdle(sessionId, fn) : fn()
const initialize = Effect.fn("ACP.initialize")(function* (params: InitializeRequest) {
const started = performance.now()
@@ -504,19 +506,21 @@ export function make(input: {
if (!command) {
const response = yield* request(
() =>
input.sdk.session.prompt(
{
sessionID: current.id,
model: {
providerID: selected.providerID,
modelID: selected.modelID,
runUntilIdle(current.id, () =>
input.sdk.session.prompt(
{
sessionID: current.id,
model: {
providerID: selected.providerID,
modelID: selected.modelID,
},
...(variant ? { variant } : {}),
parts,
...(modeId ? { agent: modeId } : {}),
directory: current.cwd,
},
...(variant ? { variant } : {}),
parts,
...(modeId ? { agent: modeId } : {}),
directory: current.cwd,
},
{ throwOnError: true },
{ throwOnError: true },
),
),
"session",
)
@@ -528,17 +532,19 @@ export function make(input: {
if (known) {
const response = yield* request(
() =>
input.sdk.session.command(
{
sessionID: current.id,
command: known.name,
arguments: command.args,
model: `${selected.providerID}/${selected.modelID}`,
...(variant ? { variant } : {}),
...(modeId ? { agent: modeId } : {}),
directory: current.cwd,
},
{ throwOnError: true },
runUntilIdle(current.id, () =>
input.sdk.session.command(
{
sessionID: current.id,
command: known.name,
arguments: command.args,
model: `${selected.providerID}/${selected.modelID}`,
...(variant ? { variant } : {}),
...(modeId ? { agent: modeId } : {}),
directory: current.cwd,
},
{ throwOnError: true },
),
),
"session",
)
@@ -549,14 +555,16 @@ export function make(input: {
if (command.name === "compact") {
yield* request(
() =>
input.sdk.session.summarize(
{
sessionID: current.id,
directory: current.cwd,
providerID: selected.providerID,
modelID: selected.modelID,
},
{ throwOnError: true },
runUntilIdle(current.id, () =>
input.sdk.session.summarize(
{
sessionID: current.id,
directory: current.cwd,
providerID: selected.providerID,
modelID: selected.modelID,
},
{ throwOnError: true },
),
),
"session",
)