Skip to content

Commit a7343e6

Browse files
committed
feat(Task): keep track of active running task
1 parent 6565b7e commit a7343e6

2 files changed

Lines changed: 145 additions & 48 deletions

File tree

src/core/task/Task.ts

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -405,6 +405,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
405405
didToolFailInCurrentTurn = false
406406
didCompleteReadingStream = false
407407
private _started = false
408+
private _runPromise: Promise<void> | undefined
408409
// No streaming parser is required.
409410
assistantMessageParser?: undefined
410411
private providerProfileChangeListener?: (config: { name: string; provider?: string }) => void
@@ -1875,6 +1876,27 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
18751876
}
18761877
}
18771878

1879+
/**
1880+
* Like `start()`, but returns the underlying promise so callers (e.g.
1881+
* `TaskScheduler`) can await task completion and gate concurrency.
1882+
* Idempotent: subsequent calls return the same in-flight promise.
1883+
*/
1884+
public run(): Promise<void> {
1885+
if (this._runPromise !== undefined) {
1886+
return this._runPromise
1887+
}
1888+
if (this._started) {
1889+
// Already launched via constructor or start() — no promise to return.
1890+
return Promise.resolve()
1891+
}
1892+
this._started = true
1893+
1894+
const { task, images } = this.metadata
1895+
1896+
this._runPromise = task || images ? this.startTask(task ?? undefined, images ?? undefined) : Promise.resolve()
1897+
return this._runPromise
1898+
}
1899+
18781900
private async startTask(task?: string, images?: string[]): Promise<void> {
18791901
try {
18801902
// `conversationHistory` (for API) and `clineMessages` (for webview)

src/core/task/__tests__/Task.dispose.test.ts

Lines changed: 123 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
import { ProviderSettings } from "@roo-code/types"
1+
import { type ProviderSettings, RooCodeEventName } from "@roo-code/types"
22

33
import { Task } from "../Task"
44
import { ClineProvider } from "../../webview/ClineProvider"
@@ -44,7 +44,11 @@ vi.mock("@roo-code/telemetry", () => ({
4444
}))
4545

4646
describe("Task dispose method", () => {
47-
let mockProvider: any
47+
let mockProvider: {
48+
context: { globalStorageUri: { fsPath: string } }
49+
getState: ReturnType<typeof vi.fn>
50+
log: ReturnType<typeof vi.fn>
51+
}
4852
let mockApiConfiguration: ProviderSettings
4953
let task: Task
5054

@@ -69,7 +73,7 @@ describe("Task dispose method", () => {
6973

7074
// Create task instance without starting it
7175
task = new Task({
72-
provider: mockProvider as ClineProvider,
76+
provider: mockProvider as unknown as ClineProvider,
7377
apiConfiguration: mockApiConfiguration,
7478
startTask: false,
7579
})
@@ -88,15 +92,14 @@ describe("Task dispose method", () => {
8892
const listener2 = vi.fn(() => {})
8993
const listener3 = vi.fn((taskId: string) => {})
9094

91-
// Use type assertion to bypass strict event typing for testing
92-
;(task as any).on("TaskStarted", listener1)
93-
;(task as any).on("TaskAborted", listener2)
94-
;(task as any).on("TaskIdle", listener3)
95+
task.on(RooCodeEventName.TaskStarted, listener1)
96+
task.on(RooCodeEventName.TaskAborted, listener2)
97+
task.on(RooCodeEventName.TaskIdle, listener3)
9598

9699
// Verify listeners are added
97-
expect(task.listenerCount("TaskStarted")).toBe(1)
98-
expect(task.listenerCount("TaskAborted")).toBe(1)
99-
expect(task.listenerCount("TaskIdle")).toBe(1)
100+
expect(task.listenerCount(RooCodeEventName.TaskStarted)).toBe(1)
101+
expect(task.listenerCount(RooCodeEventName.TaskAborted)).toBe(1)
102+
expect(task.listenerCount(RooCodeEventName.TaskIdle)).toBe(1)
100103

101104
// Spy on removeAllListeners method
102105
const removeAllListenersSpy = vi.spyOn(task, "removeAllListeners")
@@ -108,9 +111,9 @@ describe("Task dispose method", () => {
108111
expect(removeAllListenersSpy).toHaveBeenCalledOnce()
109112

110113
// Verify all listeners are removed
111-
expect(task.listenerCount("TaskStarted")).toBe(0)
112-
expect(task.listenerCount("TaskAborted")).toBe(0)
113-
expect(task.listenerCount("TaskIdle")).toBe(0)
114+
expect(task.listenerCount(RooCodeEventName.TaskStarted)).toBe(0)
115+
expect(task.listenerCount(RooCodeEventName.TaskAborted)).toBe(0)
116+
expect(task.listenerCount(RooCodeEventName.TaskIdle)).toBe(0)
114117
})
115118

116119
test("should handle errors when removing event listeners", () => {
@@ -158,53 +161,125 @@ describe("Task dispose method", () => {
158161
const listeners = {
159162
TaskStarted: vi.fn(() => {}),
160163
TaskAborted: vi.fn(() => {}),
161-
TaskIdle: vi.fn((taskId: string) => {}),
162-
TaskActive: vi.fn((taskId: string) => {}),
164+
TaskIdle: vi.fn((_taskId: string) => {}),
165+
TaskActive: vi.fn((_taskId: string) => {}),
163166
TaskAskResponded: vi.fn(() => {}),
164-
Message: vi.fn((data: { action: "created" | "updated"; message: any }) => {}),
165-
TaskTokenUsageUpdated: vi.fn((taskId: string, tokenUsage: any) => {}),
166-
TaskToolFailed: vi.fn((taskId: string, tool: any, error: string) => {}),
167-
TaskUnpaused: vi.fn(() => {}),
167+
Message: vi.fn(() => {}),
168+
TaskTokenUsageUpdated: vi.fn(() => {}),
169+
TaskToolFailed: vi.fn(() => {}),
170+
TaskUnpaused: vi.fn((_taskId: string) => {}),
168171
}
169172

170-
// Add all listeners using type assertion to bypass strict typing for testing
171-
const taskAny = task as any
172-
taskAny.on("TaskStarted", listeners.TaskStarted)
173-
taskAny.on("TaskAborted", listeners.TaskAborted)
174-
taskAny.on("TaskIdle", listeners.TaskIdle)
175-
taskAny.on("TaskActive", listeners.TaskActive)
176-
taskAny.on("TaskAskResponded", listeners.TaskAskResponded)
177-
taskAny.on("Message", listeners.Message)
178-
taskAny.on("TaskTokenUsageUpdated", listeners.TaskTokenUsageUpdated)
179-
taskAny.on("TaskToolFailed", listeners.TaskToolFailed)
180-
taskAny.on("TaskUnpaused", listeners.TaskUnpaused)
173+
task.on(RooCodeEventName.TaskStarted, listeners.TaskStarted)
174+
task.on(RooCodeEventName.TaskAborted, listeners.TaskAborted)
175+
task.on(RooCodeEventName.TaskIdle, listeners.TaskIdle)
176+
task.on(RooCodeEventName.TaskActive, listeners.TaskActive)
177+
task.on(RooCodeEventName.TaskAskResponded, listeners.TaskAskResponded)
178+
task.on(RooCodeEventName.Message, listeners.Message)
179+
task.on(RooCodeEventName.TaskTokenUsageUpdated, listeners.TaskTokenUsageUpdated)
180+
task.on(RooCodeEventName.TaskToolFailed, listeners.TaskToolFailed)
181+
task.on(RooCodeEventName.TaskUnpaused, listeners.TaskUnpaused)
181182

182183
// Verify all listeners are added
183-
expect(task.listenerCount("TaskStarted")).toBe(1)
184-
expect(task.listenerCount("TaskAborted")).toBe(1)
185-
expect(task.listenerCount("TaskIdle")).toBe(1)
186-
expect(task.listenerCount("TaskActive")).toBe(1)
187-
expect(task.listenerCount("TaskAskResponded")).toBe(1)
188-
expect(task.listenerCount("Message")).toBe(1)
189-
expect(task.listenerCount("TaskTokenUsageUpdated")).toBe(1)
190-
expect(task.listenerCount("TaskToolFailed")).toBe(1)
191-
expect(task.listenerCount("TaskUnpaused")).toBe(1)
184+
expect(task.listenerCount(RooCodeEventName.TaskStarted)).toBe(1)
185+
expect(task.listenerCount(RooCodeEventName.TaskAborted)).toBe(1)
186+
expect(task.listenerCount(RooCodeEventName.TaskIdle)).toBe(1)
187+
expect(task.listenerCount(RooCodeEventName.TaskActive)).toBe(1)
188+
expect(task.listenerCount(RooCodeEventName.TaskAskResponded)).toBe(1)
189+
expect(task.listenerCount(RooCodeEventName.Message)).toBe(1)
190+
expect(task.listenerCount(RooCodeEventName.TaskTokenUsageUpdated)).toBe(1)
191+
expect(task.listenerCount(RooCodeEventName.TaskToolFailed)).toBe(1)
192+
expect(task.listenerCount(RooCodeEventName.TaskUnpaused)).toBe(1)
192193

193194
// Call dispose
194195
task.dispose()
195196

196197
// Verify all listeners are removed
197-
expect(task.listenerCount("TaskStarted")).toBe(0)
198-
expect(task.listenerCount("TaskAborted")).toBe(0)
199-
expect(task.listenerCount("TaskIdle")).toBe(0)
200-
expect(task.listenerCount("TaskActive")).toBe(0)
201-
expect(task.listenerCount("TaskAskResponded")).toBe(0)
202-
expect(task.listenerCount("Message")).toBe(0)
203-
expect(task.listenerCount("TaskTokenUsageUpdated")).toBe(0)
204-
expect(task.listenerCount("TaskToolFailed")).toBe(0)
205-
expect(task.listenerCount("TaskUnpaused")).toBe(0)
198+
expect(task.listenerCount(RooCodeEventName.TaskStarted)).toBe(0)
199+
expect(task.listenerCount(RooCodeEventName.TaskAborted)).toBe(0)
200+
expect(task.listenerCount(RooCodeEventName.TaskIdle)).toBe(0)
201+
expect(task.listenerCount(RooCodeEventName.TaskActive)).toBe(0)
202+
expect(task.listenerCount(RooCodeEventName.TaskAskResponded)).toBe(0)
203+
expect(task.listenerCount(RooCodeEventName.Message)).toBe(0)
204+
expect(task.listenerCount(RooCodeEventName.TaskTokenUsageUpdated)).toBe(0)
205+
expect(task.listenerCount(RooCodeEventName.TaskToolFailed)).toBe(0)
206+
expect(task.listenerCount(RooCodeEventName.TaskUnpaused)).toBe(0)
206207

207208
// Verify total listener count is 0
208209
expect(task.eventNames().length).toBe(0)
209210
})
210211
})
212+
213+
describe("Task.run() idempotency", () => {
214+
// Reuses the mock setup from the outer describe block above.
215+
let mockProvider: ReturnType<typeof buildMockProvider>
216+
let mockApiConfiguration: ProviderSettings
217+
218+
function buildMockProvider() {
219+
return {
220+
context: { globalStorageUri: { fsPath: "/test/path" } },
221+
getState: vi.fn().mockResolvedValue({ mode: "code" }),
222+
log: vi.fn(),
223+
}
224+
}
225+
226+
beforeEach(() => {
227+
vi.clearAllMocks()
228+
mockProvider = buildMockProvider()
229+
mockApiConfiguration = { apiProvider: "anthropic", apiKey: "test-key" } as ProviderSettings
230+
})
231+
232+
test("run() does not invoke startTask when task was already started by constructor", async () => {
233+
// Spy on the prototype before construction so we capture the constructor's call too.
234+
const startTaskSpy = vi.spyOn(Task.prototype as any, "startTask").mockResolvedValue(undefined)
235+
236+
const t = new Task({
237+
provider: mockProvider as unknown as ClineProvider,
238+
apiConfiguration: mockApiConfiguration,
239+
task: "hello",
240+
startTask: true,
241+
})
242+
243+
const callsBefore = startTaskSpy.mock.calls.length // constructor fired it once
244+
void t.run()
245+
expect(startTaskSpy.mock.calls.length).toBe(callsBefore) // run() must not add a second call
246+
t.dispose()
247+
startTaskSpy.mockRestore()
248+
})
249+
250+
test("run() does not invoke startTask when task was already started by start()", async () => {
251+
const startTaskSpy = vi.spyOn(Task.prototype as any, "startTask").mockResolvedValue(undefined)
252+
253+
const t = new Task({
254+
provider: mockProvider as unknown as ClineProvider,
255+
apiConfiguration: mockApiConfiguration,
256+
task: "hello",
257+
startTask: false,
258+
})
259+
t.start()
260+
const callsAfterStart = startTaskSpy.mock.calls.length // start() fired it once
261+
262+
void t.run()
263+
expect(startTaskSpy.mock.calls.length).toBe(callsAfterStart) // no additional call
264+
t.dispose()
265+
startTaskSpy.mockRestore()
266+
})
267+
268+
test("run() returns the same promise on repeated calls", async () => {
269+
const startTaskSpy = vi.spyOn(Task.prototype as any, "startTask").mockResolvedValue(undefined)
270+
271+
const t = new Task({
272+
provider: mockProvider as unknown as ClineProvider,
273+
apiConfiguration: mockApiConfiguration,
274+
task: "hello",
275+
startTask: false,
276+
})
277+
278+
const p1 = t.run()
279+
const p2 = t.run()
280+
expect(p1).toBe(p2)
281+
await p1
282+
t.dispose()
283+
startTaskSpy.mockRestore()
284+
})
285+
})

0 commit comments

Comments
 (0)