fix(acp): drain updates before end turn (#40422)
This commit is contained in:
@@ -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"
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user