From 130907587e0178ee8cab9508116aa99018b4ecb9 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 9 Aug 2026 14:54:43 +0000 Subject: [PATCH 1/8] FB-15a: extract Task steering orchestration to TaskSteering.ts MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Moves claim/commit/rollback/settle/enqueue/append/restore helpers out of src/core/task/index.ts into a focused helper module. - Keeps Task class as a thin façade that delegates via a context object, preserving the existing public/private method surface. - Verified: core/task tests 283 passing, tsc baseline unchanged, biome clean. Co-Authored-By: Alexandros Salapatas --- src/core/task/TaskSteering.ts | 266 ++++++++++++++++++++++++++++++++++ src/core/task/index.ts | 265 +++++---------------------------- 2 files changed, 305 insertions(+), 226 deletions(-) create mode 100644 src/core/task/TaskSteering.ts diff --git a/src/core/task/TaskSteering.ts b/src/core/task/TaskSteering.ts new file mode 100644 index 00000000..92534b18 --- /dev/null +++ b/src/core/task/TaskSteering.ts @@ -0,0 +1,266 @@ +import { findSlashCommandInTags } from "@core/slash-commands/commandParser" +import type { DiracMessage } from "@shared/ExtensionMessage" +import { DiracMessageType, SteeringTranscriptStatus, TaskStatus } from "@shared/ExtensionMessage" +import type { DiracContent, DiracStorageMessage, DiracTextContentBlock } from "@shared/messages/content" +import { ulid } from "ulid" +import type { MessageStateHandler } from "./message-state" +import { + collectDeliveredSteeringMessageIds, + formatSteeringMessages, + restoreQueuedSteeringMessages, + SteeringClaim, + SteeringDeliveryState, + SteeringMessage, +} from "./steering" +import type { TaskMessenger } from "./TaskMessenger" +import type { TaskState } from "./TaskState" + +export interface TaskSteeringContext { + taskState: TaskState + messageStateHandler: MessageStateHandler + taskMessenger: TaskMessenger + postStateToWebview: () => void | Promise + withStateLock: (fn: () => T | Promise) => Promise +} + +export function canAcceptSteeringMessage(ctx: TaskSteeringContext): boolean { + if (ctx.taskState.completionCommitted) return false + if (ctx.taskState.abort || ctx.taskState.pendingTaskReplacement) return false + if (ctx.taskState.waitingCardIds.length > 0) return false + return ![TaskStatus.IDLE, TaskStatus.COMPLETED, TaskStatus.CANCELLED, TaskStatus.CANCELLING].includes(ctx.taskState.status) +} + +export async function enqueueSteeringMessage(ctx: TaskSteeringContext, text: string): Promise { + const normalizedText = text.trim() + if (!normalizedText) throw new Error("Steering guidance cannot be empty") + + const steeringMessage = await ctx.withStateLock(async (): Promise => { + if (!canAcceptSteeringMessage(ctx)) { + throw new Error(`Task cannot accept steering while ${ctx.taskState.status}`) + } + + const createdAt = Date.now() + const transcriptMessageId = ctx.taskMessenger.generateId() + const message: SteeringMessage = { + id: ulid(), + text: normalizedText, + createdAt, + transcriptMessageId, + deliveryState: SteeringDeliveryState.QUEUED, + } + const transcriptMessage: DiracMessage = { + id: transcriptMessageId, + ts: createdAt, + content: { + type: DiracMessageType.MARKDOWN, + content: normalizedText, + role: "user", + steering: { status: SteeringTranscriptStatus.QUEUED }, + }, + } + await ctx.messageStateHandler.addToDiracMessages(transcriptMessage) + ctx.taskState.steeringMessages.push(message) + return message + }) + + await ctx.postStateToWebview() + return steeringMessage.transcriptMessageId +} + +export async function claimSteeringMessages(ctx: TaskSteeringContext): Promise { + return ctx.withStateLock(() => { + const messages = ctx.taskState.steeringMessages.filter( + (message) => message.deliveryState === SteeringDeliveryState.QUEUED, + ) + if (messages.length === 0) return undefined + + const claimId = ulid() + for (const message of messages) { + message.deliveryState = SteeringDeliveryState.CLAIMED + message.claimId = claimId + } + return { id: claimId, messages: messages.map((message) => ({ ...message })) } + }) +} + +export async function commitSteeringClaim(ctx: TaskSteeringContext, claimId: string): Promise { + const claimedMessages = await ctx.withStateLock(() => { + const messages = ctx.taskState.steeringMessages.filter( + (message) => message.deliveryState === SteeringDeliveryState.CLAIMED && message.claimId === claimId, + ) + for (const message of messages) { + message.deliveryState = SteeringDeliveryState.SENT + message.claimId = undefined + } + return messages.map((message) => ({ ...message })) + }) + + const transcriptMessages = claimedMessages.map((message) => { + const index = ctx.messageStateHandler.findMessageIndexById(message.transcriptMessageId) + if (index === -1) throw new Error(`Steering transcript message not found: ${message.transcriptMessageId}`) + const transcriptMessage = ctx.messageStateHandler.getDiracMessages()[index] + if (transcriptMessage.content.type !== DiracMessageType.MARKDOWN) { + throw new Error(`Steering transcript message is not markdown: ${message.transcriptMessageId}`) + } + return { index, content: transcriptMessage.content } + }) + + for (const transcript of transcriptMessages) { + await ctx.messageStateHandler.updateDiracMessage(transcript.index, { + content: { + ...transcript.content, + steering: { status: SteeringTranscriptStatus.SENT }, + }, + }) + } + await ctx.messageStateHandler.saveDiracMessagesAndUpdateHistory() + await ctx.postStateToWebview() +} + +export async function settleConsumedSteeringClaim(ctx: TaskSteeringContext, claim: SteeringClaim): Promise { + let receiptError: unknown + try { + await ctx.messageStateHandler.recordDeliveredSteeringMessageIds( + claim.messages.map((message) => message.transcriptMessageId), + ) + } catch (error) { + receiptError = error + } + + try { + await commitSteeringClaim(ctx, claim.id) + } catch (commitError) { + if (receiptError) { + throw new AggregateError([receiptError, commitError], "Failed to persist consumed steering delivery") + } + throw commitError + } + if (receiptError) throw receiptError +} + +export async function rollbackSteeringClaim(ctx: TaskSteeringContext, claimId: string): Promise { + await ctx.withStateLock(() => { + for (const message of ctx.taskState.steeringMessages) { + if (message.deliveryState !== SteeringDeliveryState.CLAIMED || message.claimId !== claimId) continue + message.deliveryState = SteeringDeliveryState.QUEUED + message.claimId = undefined + } + }) +} + +export function isSlashCommandSteeringMessage(message: Pick): boolean { + return findSlashCommandInTags(formatSteeringMessages([message])) !== null +} + +export async function releaseSteeringClaimSuffix( + ctx: TaskSteeringContext, + claim: SteeringClaim, + retainedCount: number, +): Promise { + const retainedMessages = claim.messages.slice(0, retainedCount) + const retainedMessageIds = new Set(retainedMessages.map((message) => message.id)) + await ctx.withStateLock(() => { + for (const message of ctx.taskState.steeringMessages) { + if (message.deliveryState !== SteeringDeliveryState.CLAIMED || message.claimId !== claim.id) continue + if (retainedMessageIds.has(message.id)) continue + message.deliveryState = SteeringDeliveryState.QUEUED + message.claimId = undefined + } + }) + return { ...claim, messages: retainedMessages } +} + +export async function appendQueuedSteeringToUserContent( + ctx: TaskSteeringContext, + userContent: DiracContent[], +): Promise { + const steeringClaim = await claimSteeringMessages(ctx) + if (!steeringClaim) return undefined + const commandIndex = steeringClaim.messages.findIndex((message) => isSlashCommandSteeringMessage(message)) + if (commandIndex === -1) { + userContent.push({ + type: "text", + text: formatSteeringMessages(steeringClaim.messages), + isUserInput: true, + steeringMessageIds: steeringClaim.messages.map((message) => message.transcriptMessageId), + }) + return steeringClaim + } + + const retainedCount = commandIndex === 0 ? 1 : commandIndex + const requestClaim = await releaseSteeringClaimSuffix(ctx, steeringClaim, retainedCount) + userContent.push({ + type: "text", + text: formatSteeringMessages(requestClaim.messages), + isUserInput: true, + steeringMessageIds: requestClaim.messages.map((message) => message.transcriptMessageId), + }) + return requestClaim +} + +export async function appendQueuedSteeringToNextApiRequest( + ctx: TaskSteeringContext, + outboundHistory: DiracStorageMessage[], +): Promise { + const steeringClaim = await claimSteeringMessages(ctx) + if (!steeringClaim) return + if (steeringClaim.messages.some((message) => isSlashCommandSteeringMessage(message))) { + await rollbackSteeringClaim(ctx, steeringClaim.id) + return + } + + const messageIds = steeringClaim.messages.map((message) => message.transcriptMessageId) + const steeringBlock: DiracTextContentBlock = { + type: "text", + text: formatSteeringMessages(steeringClaim.messages), + steeringMessageIds: messageIds, + } + const outboundUserMessage = outboundHistory.at(-1) + if (!outboundUserMessage || outboundUserMessage.role !== "user") { + await rollbackSteeringClaim(ctx, steeringClaim.id) + throw new Error("Cannot steer a model request without a final user message") + } + + let userMessagePersisted = false + try { + const persistedUserMessage = await ctx.messageStateHandler.appendToLastApiConversationUserMessage(steeringBlock) + userMessagePersisted = true + if (outboundUserMessage !== persistedUserMessage) { + if (typeof outboundUserMessage.content === "string") { + outboundUserMessage.content = [{ type: "text", text: outboundUserMessage.content }, steeringBlock] + } else { + outboundUserMessage.content.push(steeringBlock) + } + } + await settleConsumedSteeringClaim(ctx, steeringClaim) + } catch (error) { + if (!userMessagePersisted) await rollbackSteeringClaim(ctx, steeringClaim.id) + throw error + } +} + +export async function commitAttemptCompletion(ctx: TaskSteeringContext): Promise { + return ctx.withStateLock(() => { + const hasQueuedSteering = ctx.taskState.steeringMessages.some( + (message) => message.deliveryState === SteeringDeliveryState.QUEUED, + ) + if (hasQueuedSteering) { + ctx.taskState.didAttemptCompletion = false + return false + } + ctx.taskState.completionCommitted = true + ctx.taskState.didAttemptCompletion = true + return true + }) +} + +export function restoreQueuedSteeringFromTranscript(ctx: TaskSteeringContext): void { + const deliveredMessageIds = collectDeliveredSteeringMessageIds(ctx.messageStateHandler.getApiConversationHistory()) + for (const messageId of ctx.messageStateHandler.getApiConversationProviderState().deliveredSteeringMessageIds ?? []) { + deliveredMessageIds.add(messageId) + } + ctx.taskState.steeringMessages = restoreQueuedSteeringMessages( + ctx.messageStateHandler.getDiracMessages(), + deliveredMessageIds, + ) +} diff --git a/src/core/task/index.ts b/src/core/task/index.ts index e62222fd..dbf3479d 100644 --- a/src/core/task/index.ts +++ b/src/core/task/index.ts @@ -23,7 +23,6 @@ import { CommandPermissionController } from "@core/permissions" import type { SystemPromptContext } from "@core/prompts/system-prompt" import { getSystemPrompt } from "@core/prompts/system-prompt" import type { SlashCommandDirectAction } from "@core/slash-commands" -import { findSlashCommandInTags } from "@core/slash-commands/commandParser" import { ensureRulesDirectoryExists, ensureTaskDirectoryExists } from "@core/storage/disk" import { createDefaultTextCondensationTemplateRegistry, TASK_HANDOFF_TEMPLATE_ID } from "@core/text-condensation/templates" import { isUtilityTextCondensationAvailable } from "@core/text-condensation/UtilityTextCondensationAvailability" @@ -55,15 +54,7 @@ import { ApiConfiguration } from "@shared/api" import { findLastIndex } from "@shared/array" import { DiracClient } from "@shared/dirac" import { getExtensionSourceDir } from "@shared/dirac/constants" -import { - CardStatus, - DiracApiReqCancelReason, - DiracMessage, - DiracMessageContent, - DiracMessageType, - SteeringTranscriptStatus, - TaskStatus, -} from "@shared/ExtensionMessage" +import { CardStatus, DiracApiReqCancelReason, DiracMessageContent, DiracMessageType, TaskStatus } from "@shared/ExtensionMessage" import { HistoryItem } from "@shared/HistoryItem" import { DEFAULT_LANGUAGE_SETTINGS, getLanguageKey, LanguageDisplay } from "@shared/Languages" import { @@ -108,16 +99,22 @@ import { ResponseProcessor } from "./ResponseProcessor" import { StreamChunkCoordinator } from "./StreamChunkCoordinator" import { StreamingMetricsManager } from "./StreamingMetricsManager" import { StreamResponseHandler } from "./StreamResponseHandler" -import { - collectDeliveredSteeringMessageIds, - formatSteeringMessages, - restoreQueuedSteeringMessages, - type SteeringClaim, - SteeringDeliveryState, - type SteeringMessage, -} from "./steering" +import { type SteeringClaim } from "./steering" import { TaskMessenger } from "./TaskMessenger" import { TaskState } from "./TaskState" +import { + appendQueuedSteeringToNextApiRequest, + appendQueuedSteeringToUserContent, + canAcceptSteeringMessage, + claimSteeringMessages, + commitAttemptCompletion, + commitSteeringClaim, + enqueueSteeringMessage, + restoreQueuedSteeringFromTranscript, + rollbackSteeringClaim, + settleConsumedSteeringClaim, + type TaskSteeringContext, +} from "./TaskSteering" import { ToolExecutor } from "./ToolExecutor" import { DiracContext } from "./tools/context/DiracContext" import type { ToolSnapshotDirtyReason } from "./tools/runtime/ToolSnapshot" @@ -175,227 +172,50 @@ export class Task { return await this.stateMutex.withLock(fn) } + private get steeringContext(): TaskSteeringContext { + return { + taskState: this.taskState, + messageStateHandler: this.messageStateHandler, + taskMessenger: this.taskMessenger, + postStateToWebview: () => this.postStateToWebview(), + withStateLock: (fn) => this.withStateLock(fn), + } + } + public canAcceptSteeringMessage(): boolean { - if (this.taskState.completionCommitted) return false - if (this.taskState.abort || this.taskState.pendingTaskReplacement) return false - if (this.taskState.waitingCardIds.length > 0) return false - return ![TaskStatus.IDLE, TaskStatus.COMPLETED, TaskStatus.CANCELLED, TaskStatus.CANCELLING].includes( - this.taskState.status, - ) + return canAcceptSteeringMessage(this.steeringContext) } public async enqueueSteeringMessage(text: string): Promise { - const normalizedText = text.trim() - if (!normalizedText) throw new Error("Steering guidance cannot be empty") - - const steeringMessage = await this.withStateLock(async (): Promise => { - if (!this.canAcceptSteeringMessage()) { - throw new Error(`Task cannot accept steering while ${this.taskState.status}`) - } - - const createdAt = Date.now() - const transcriptMessageId = this.taskMessenger.generateId() - const message: SteeringMessage = { - id: ulid(), - text: normalizedText, - createdAt, - transcriptMessageId, - deliveryState: SteeringDeliveryState.QUEUED, - } - const transcriptMessage: DiracMessage = { - id: transcriptMessageId, - ts: createdAt, - content: { - type: DiracMessageType.MARKDOWN, - content: normalizedText, - role: "user", - steering: { status: SteeringTranscriptStatus.QUEUED }, - }, - } - await this.messageStateHandler.addToDiracMessages(transcriptMessage) - this.taskState.steeringMessages.push(message) - return message - }) - - await this.postStateToWebview() - return steeringMessage.transcriptMessageId + return enqueueSteeringMessage(this.steeringContext, text) } private async claimSteeringMessages(): Promise { - return this.withStateLock(() => { - const messages = this.taskState.steeringMessages.filter( - (message) => message.deliveryState === SteeringDeliveryState.QUEUED, - ) - if (messages.length === 0) return undefined - - const claimId = ulid() - for (const message of messages) { - message.deliveryState = SteeringDeliveryState.CLAIMED - message.claimId = claimId - } - return { id: claimId, messages: messages.map((message) => ({ ...message })) } - }) + return claimSteeringMessages(this.steeringContext) } private async commitSteeringClaim(claimId: string): Promise { - const claimedMessages = await this.withStateLock(() => { - const messages = this.taskState.steeringMessages.filter( - (message) => message.deliveryState === SteeringDeliveryState.CLAIMED && message.claimId === claimId, - ) - for (const message of messages) { - message.deliveryState = SteeringDeliveryState.SENT - message.claimId = undefined - } - return messages.map((message) => ({ ...message })) - }) - - const transcriptMessages = claimedMessages.map((message) => { - const index = this.messageStateHandler.findMessageIndexById(message.transcriptMessageId) - if (index === -1) throw new Error(`Steering transcript message not found: ${message.transcriptMessageId}`) - const transcriptMessage = this.messageStateHandler.getDiracMessages()[index] - if (transcriptMessage.content.type !== DiracMessageType.MARKDOWN) { - throw new Error(`Steering transcript message is not markdown: ${message.transcriptMessageId}`) - } - return { index, content: transcriptMessage.content } - }) - - for (const transcript of transcriptMessages) { - await this.messageStateHandler.updateDiracMessage(transcript.index, { - content: { - ...transcript.content, - steering: { status: SteeringTranscriptStatus.SENT }, - }, - }) - } - await this.messageStateHandler.saveDiracMessagesAndUpdateHistory() - await this.postStateToWebview() + return commitSteeringClaim(this.steeringContext, claimId) } private async settleConsumedSteeringClaim(claim: SteeringClaim): Promise { - let receiptError: unknown - try { - await this.messageStateHandler.recordDeliveredSteeringMessageIds( - claim.messages.map((message) => message.transcriptMessageId), - ) - } catch (error) { - receiptError = error - } - - try { - await this.commitSteeringClaim(claim.id) - } catch (commitError) { - if (receiptError) { - throw new AggregateError([receiptError, commitError], "Failed to persist consumed steering delivery") - } - throw commitError - } - if (receiptError) throw receiptError + return settleConsumedSteeringClaim(this.steeringContext, claim) } private async rollbackSteeringClaim(claimId: string): Promise { - await this.withStateLock(() => { - for (const message of this.taskState.steeringMessages) { - if (message.deliveryState !== SteeringDeliveryState.CLAIMED || message.claimId !== claimId) continue - message.deliveryState = SteeringDeliveryState.QUEUED - message.claimId = undefined - } - }) - } - - private isSlashCommandSteeringMessage(message: Pick): boolean { - return findSlashCommandInTags(formatSteeringMessages([message])) !== null - } - - private async releaseSteeringClaimSuffix(claim: SteeringClaim, retainedCount: number): Promise { - const retainedMessages = claim.messages.slice(0, retainedCount) - const retainedMessageIds = new Set(retainedMessages.map((message) => message.id)) - await this.withStateLock(() => { - for (const message of this.taskState.steeringMessages) { - if (message.deliveryState !== SteeringDeliveryState.CLAIMED || message.claimId !== claim.id) continue - if (retainedMessageIds.has(message.id)) continue - message.deliveryState = SteeringDeliveryState.QUEUED - message.claimId = undefined - } - }) - return { ...claim, messages: retainedMessages } + return rollbackSteeringClaim(this.steeringContext, claimId) } private async appendQueuedSteeringToUserContent(userContent: DiracContent[]): Promise { - const steeringClaim = await this.claimSteeringMessages() - if (!steeringClaim) return undefined - const commandIndex = steeringClaim.messages.findIndex((message) => this.isSlashCommandSteeringMessage(message)) - if (commandIndex === -1) { - userContent.push({ - type: "text", - text: formatSteeringMessages(steeringClaim.messages), - isUserInput: true, - steeringMessageIds: steeringClaim.messages.map((message) => message.transcriptMessageId), - }) - return steeringClaim - } - - const retainedCount = commandIndex === 0 ? 1 : commandIndex - const requestClaim = await this.releaseSteeringClaimSuffix(steeringClaim, retainedCount) - userContent.push({ - type: "text", - text: formatSteeringMessages(requestClaim.messages), - isUserInput: true, - steeringMessageIds: requestClaim.messages.map((message) => message.transcriptMessageId), - }) - return requestClaim + return appendQueuedSteeringToUserContent(this.steeringContext, userContent) } private async appendQueuedSteeringToNextApiRequest(outboundHistory: DiracStorageMessage[]): Promise { - const steeringClaim = await this.claimSteeringMessages() - if (!steeringClaim) return - if (steeringClaim.messages.some((message) => this.isSlashCommandSteeringMessage(message))) { - await this.rollbackSteeringClaim(steeringClaim.id) - return - } - - const messageIds = steeringClaim.messages.map((message) => message.transcriptMessageId) - const steeringBlock: DiracTextContentBlock = { - type: "text", - text: formatSteeringMessages(steeringClaim.messages), - steeringMessageIds: messageIds, - } - const outboundUserMessage = outboundHistory.at(-1) - if (!outboundUserMessage || outboundUserMessage.role !== "user") { - await this.rollbackSteeringClaim(steeringClaim.id) - throw new Error("Cannot steer a model request without a final user message") - } - - let userMessagePersisted = false - try { - const persistedUserMessage = await this.messageStateHandler.appendToLastApiConversationUserMessage(steeringBlock) - userMessagePersisted = true - if (outboundUserMessage !== persistedUserMessage) { - if (typeof outboundUserMessage.content === "string") { - outboundUserMessage.content = [{ type: "text", text: outboundUserMessage.content }, steeringBlock] - } else { - outboundUserMessage.content.push(steeringBlock) - } - } - await this.settleConsumedSteeringClaim(steeringClaim) - } catch (error) { - if (!userMessagePersisted) await this.rollbackSteeringClaim(steeringClaim.id) - throw error - } + return appendQueuedSteeringToNextApiRequest(this.steeringContext, outboundHistory) } private async commitAttemptCompletion(): Promise { - return this.withStateLock(() => { - const hasQueuedSteering = this.taskState.steeringMessages.some( - (message) => message.deliveryState === SteeringDeliveryState.QUEUED, - ) - if (hasQueuedSteering) { - this.taskState.didAttemptCompletion = false - return false - } - this.taskState.completionCommitted = true - this.taskState.didAttemptCompletion = true - return true - }) + return commitAttemptCompletion(this.steeringContext) } public async setActiveHookExecution(hookExecution: NonNullable): Promise { @@ -1221,14 +1041,7 @@ export class Task { } public restoreQueuedSteeringFromTranscript(): void { - const deliveredMessageIds = collectDeliveredSteeringMessageIds(this.messageStateHandler.getApiConversationHistory()) - for (const messageId of this.messageStateHandler.getApiConversationProviderState().deliveredSteeringMessageIds ?? []) { - deliveredMessageIds.add(messageId) - } - this.taskState.steeringMessages = restoreQueuedSteeringMessages( - this.messageStateHandler.getDiracMessages(), - deliveredMessageIds, - ) + restoreQueuedSteeringFromTranscript(this.steeringContext) } public async resumeTaskFromHistory() { @@ -2121,19 +1934,19 @@ export class Task { afterUserContentPersisted: async () => { steeringClaimConsumed = true if (!steeringClaim) return - await this.settleConsumedSteeringClaim(steeringClaim) + await settleConsumedSteeringClaim(this.steeringContext, steeringClaim) }, }) if (steeringClaim && !steeringClaimConsumed) { if (apiRequestData.didConsumeUserContent) { steeringClaimConsumed = true - await this.settleConsumedSteeringClaim(steeringClaim) + await settleConsumedSteeringClaim(this.steeringContext, steeringClaim) } else { - await this.rollbackSteeringClaim(steeringClaim.id) + await rollbackSteeringClaim(this.steeringContext, steeringClaim.id) } } } catch (error) { - if (steeringClaim && !steeringClaimConsumed) await this.rollbackSteeringClaim(steeringClaim.id) + if (steeringClaim && !steeringClaimConsumed) await rollbackSteeringClaim(this.steeringContext, steeringClaim.id) throw error } userContent = apiRequestData.userContent From a2ef20c0ad179a055fd207f71e7e921d3ea9548b Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 9 Aug 2026 14:57:26 +0000 Subject: [PATCH 2/8] FB-15a: extract prompt metadata artifact writer to TaskPromptArtifacts.ts MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Moves the ~120-line writePromptMetadataArtifacts implementation out of src/core/task/index.ts into a side-effect-isolated helper module. - Task.writePromptMetadataArtifacts becomes a thin façade; DiracContext receives an arrow wrapper so the public callback shape is preserved. - Verified: core/task tests 283 passing, tsc baseline unchanged, biome clean. Co-Authored-By: Alexandros Salapatas --- src/core/task/TaskPromptArtifacts.ts | 125 +++++++++++++++++++++++++++ src/core/task/index.ts | 124 +++----------------------- 2 files changed, 137 insertions(+), 112 deletions(-) create mode 100644 src/core/task/TaskPromptArtifacts.ts diff --git a/src/core/task/TaskPromptArtifacts.ts b/src/core/task/TaskPromptArtifacts.ts new file mode 100644 index 00000000..f6925962 --- /dev/null +++ b/src/core/task/TaskPromptArtifacts.ts @@ -0,0 +1,125 @@ +import { Logger } from "@shared/services/Logger" +import fs from "fs/promises" +import * as path from "path" +import type { StateManager } from "../storage/StateManager" + +export interface TaskPromptArtifactsContext { + taskId: string + cwd: string + stateManager: StateManager +} + +export async function writePromptMetadataArtifacts( + ctx: TaskPromptArtifactsContext, + params: { + systemPrompt: string + providerInfo: { providerId?: string; modelId?: string } + tools?: any[] + fullHistory?: any[] + deletedRange?: [number, number] + }, +): Promise { + const enabledSetting = ctx.stateManager.getGlobalSettingsKey("writePromptMetadataEnabled") + const enabledFlag = process.env.DIRAC_WRITE_PROMPT_ARTIFACTS?.toLowerCase() + const enabled = + enabledSetting || enabledFlag === "1" || enabledFlag === "true" || enabledFlag === "yes" || process.env.IS_DEV === "true" + if (!enabled) { + return + } + + try { + // Env var is OS-level (user-controlled, safe to allow absolute); workspace setting is the exfiltration vector. + const envDir = process.env.DIRAC_PROMPT_ARTIFACT_DIR?.trim() + const settingDir = ctx.stateManager.getGlobalSettingsKey("writePromptMetadataDirectory")?.trim() + const cwdResolved = path.resolve(ctx.cwd) + // Setting-configured dirs must resolve under cwd to prevent workspace settings from exfiltrating prompts. + // Only validate the setting when no env var is provided — env takes precedence and is trusted. + if (!envDir && settingDir) { + const resolved = path.isAbsolute(settingDir) ? path.resolve(settingDir) : path.resolve(ctx.cwd, settingDir) + if (resolved !== cwdResolved && !resolved.startsWith(cwdResolved + path.sep)) { + Logger.warn(`[Task ${ctx.taskId}] writePromptMetadataDirectory outside cwd rejected: ${resolved}`) + return + } + } + const configuredDir = envDir || settingDir + const artifactDir = configuredDir + ? path.isAbsolute(configuredDir) + ? path.resolve(configuredDir) + : path.resolve(ctx.cwd, configuredDir) + : path.resolve(ctx.cwd, ".dirac-prompt-artifacts") + + await fs.mkdir(artifactDir, { recursive: true }) + // Defense-in-depth: re-check the boundary after mkdir resolves any symlinks in the path, + // so a workspace-planted symlink can't exfiltrate prompts outside cwd. + // Only enforced for setting-derived paths — env vars are OS-level and may legitimately point anywhere. + let writeDir = artifactDir + if (!envDir) { + const realArtifactDir = await fs.realpath(artifactDir) + if (realArtifactDir !== cwdResolved && !realArtifactDir.startsWith(cwdResolved + path.sep)) { + Logger.warn(`[Task ${ctx.taskId}] artifact dir resolves outside cwd (symlink?), rejected: ${realArtifactDir}`) + return + } + writeDir = realArtifactDir + } + // Ensure the artifact dir is git-ignored so debug dumps don't get committed. + const gitignorePath = path.join(writeDir, ".gitignore") + await fs.writeFile(gitignorePath, "*\n!.gitignore\n", "utf8").catch(() => {}) + + const debugPath = path.join(writeDir, `task-${ctx.taskId}-debug.md`) + + let markdown = `## System Prompt\n\n${params.systemPrompt}\n\n` + + if (params.tools) { + markdown += `## Tools\n\n\`\`\`json\n${JSON.stringify(params.tools, null, 2)}\n\`\`\`\n\n` + } + + if (params.fullHistory) { + markdown += `## Conversation History\n\n` + const [deletedStart, deletedEnd] = params.deletedRange || [-1, -1] + + for (let i = 0; i < params.fullHistory.length; i++) { + const message = params.fullHistory[i] + const isTruncated = i >= deletedStart && i <= deletedEnd + + markdown += `### [${message.role.toUpperCase()}]${isTruncated ? " [TRUNCATED]" : ""}\n` + + if (typeof message.content === "string") { + markdown += `${message.content}\n\n` + } else if (Array.isArray(message.content)) { + for (const block of message.content) { + if (block.type === "text") { + markdown += `**Text:** ${block.call_id ? `(\`call_id: ${block.call_id}\`)` : ""}\n${block.text}\n\n` + } else if (block.type === "thinking") { + markdown += `**Thinking:** ${block.call_id ? `(\`call_id: ${block.call_id}\`)` : ""}\n${block.thinking}\n\n` + } else if (block.type === "redacted_thinking") { + markdown += `**Thinking:** [Redacted] ${block.call_id ? `(\`call_id: ${block.call_id}\`)` : ""}\n\n` + } else if (block.type === "tool_use") { + markdown += `**Tool Use:** \`${block.name}\` (\`id: ${block.id}\`, \`call_id: ${block.call_id}\`)\n` + markdown += `\`\`\`json\n${JSON.stringify(block.input, null, 2)}\n\`\`\`\n\n` + } else if (block.type === "tool_result") { + markdown += `**Tool Result:** (\`${block.tool_use_id}\`)\n` + if (typeof block.content === "string") { + markdown += `${block.content}\n\n` + } else if (Array.isArray(block.content)) { + for (const contentBlock of block.content) { + if (contentBlock.type === "text") { + markdown += `${contentBlock.text}\n\n` + } else if (contentBlock.type === "image") { + markdown += `[Image: ${contentBlock.source?.type}]\n\n` + } + } + } + } else if (block.type === "image") { + markdown += `[Image: ${block.source?.type}]\n\n` + } + } + } + markdown += "---\n\n" + } + } + + await fs.writeFile(debugPath, markdown, "utf8") + } catch (error) { + Logger.error("Failed to write prompt metadata artifacts:", error) + } +} diff --git a/src/core/task/index.ts b/src/core/task/index.ts index dbf3479d..be521948 100644 --- a/src/core/task/index.ts +++ b/src/core/task/index.ts @@ -73,7 +73,6 @@ import { type Mode } from "@shared/storage/types" import { isMutatingTool } from "@shared/tools" import { DiracAskResponse } from "@shared/WebviewMessage" import { isLocalModel, isParallelToolCallingEnabled } from "@utils/model-utils" -import fs from "fs/promises" import Mutex from "p-mutex" import pWaitFor from "p-wait-for" import * as path from "path" @@ -101,6 +100,7 @@ import { StreamingMetricsManager } from "./StreamingMetricsManager" import { StreamResponseHandler } from "./StreamResponseHandler" import { type SteeringClaim } from "./steering" import { TaskMessenger } from "./TaskMessenger" +import { type TaskPromptArtifactsContext, writePromptMetadataArtifacts } from "./TaskPromptArtifacts" import { TaskState } from "./TaskState" import { appendQueuedSteeringToNextApiRequest, @@ -182,6 +182,14 @@ export class Task { } } + private get promptArtifactsContext(): TaskPromptArtifactsContext { + return { + taskId: this.taskId, + cwd: this.cwd, + stateManager: this.stateManager, + } + } + public canAcceptSteeringMessage(): boolean { return canAcceptSteeringMessage(this.steeringContext) } @@ -707,7 +715,7 @@ export class Task { getCurrentProviderInfo: this.getCurrentProviderInfo.bind(this), getEnvironmentDetails: this.getEnvironmentDetails.bind(this), getPinnedContext: () => this.taskState.pinnedContext, - writePromptMetadataArtifacts: this.writePromptMetadataArtifacts.bind(this), + writePromptMetadataArtifacts: (params) => writePromptMetadataArtifacts(this.promptArtifactsContext, params), handleHookCancellation: this.hookManager.handleHookCancellation.bind(this.hookManager), setActiveHookExecution: this.hookManager.setActiveHookExecution.bind(this.hookManager), clearActiveHookExecution: this.hookManager.clearActiveHookExecution.bind(this.hookManager), @@ -1140,115 +1148,7 @@ export class Task { fullHistory?: any[] deletedRange?: [number, number] }): Promise { - const enabledSetting = this.stateManager.getGlobalSettingsKey("writePromptMetadataEnabled") - const enabledFlag = process.env.DIRAC_WRITE_PROMPT_ARTIFACTS?.toLowerCase() - const enabled = - enabledSetting || - enabledFlag === "1" || - enabledFlag === "true" || - enabledFlag === "yes" || - process.env.IS_DEV === "true" - if (!enabled) { - return - } - - try { - // Env var is OS-level (user-controlled, safe to allow absolute); workspace setting is the exfiltration vector. - const envDir = process.env.DIRAC_PROMPT_ARTIFACT_DIR?.trim() - const settingDir = this.stateManager.getGlobalSettingsKey("writePromptMetadataDirectory")?.trim() - const cwdResolved = path.resolve(this.cwd) - // Setting-configured dirs must resolve under cwd to prevent workspace settings from exfiltrating prompts. - // Only validate the setting when no env var is provided — env takes precedence and is trusted. - if (!envDir && settingDir) { - const resolved = path.isAbsolute(settingDir) ? path.resolve(settingDir) : path.resolve(this.cwd, settingDir) - if (resolved !== cwdResolved && !resolved.startsWith(cwdResolved + path.sep)) { - Logger.warn(`[Task ${this.taskId}] writePromptMetadataDirectory outside cwd rejected: ${resolved}`) - return - } - } - const configuredDir = envDir || settingDir - const artifactDir = configuredDir - ? path.isAbsolute(configuredDir) - ? path.resolve(configuredDir) - : path.resolve(this.cwd, configuredDir) - : path.resolve(this.cwd, ".dirac-prompt-artifacts") - - await fs.mkdir(artifactDir, { recursive: true }) - // Defense-in-depth: re-check the boundary after mkdir resolves any symlinks in the path, - // so a workspace-planted symlink can't exfiltrate prompts outside cwd. - // Only enforced for setting-derived paths — env vars are OS-level and may legitimately point anywhere. - let writeDir = artifactDir - if (!envDir) { - const realArtifactDir = await fs.realpath(artifactDir) - if (realArtifactDir !== cwdResolved && !realArtifactDir.startsWith(cwdResolved + path.sep)) { - Logger.warn( - `[Task ${this.taskId}] artifact dir resolves outside cwd (symlink?), rejected: ${realArtifactDir}`, - ) - return - } - writeDir = realArtifactDir - } - // Ensure the artifact dir is git-ignored so debug dumps don't get committed. - const gitignorePath = path.join(writeDir, ".gitignore") - await fs.writeFile(gitignorePath, "*\n!.gitignore\n", "utf8").catch(() => {}) - - const debugPath = path.join(writeDir, `task-${this.taskId}-debug.md`) - - let markdown = `## System Prompt\n\n${params.systemPrompt}\n\n` - - if (params.tools) { - markdown += `## Tools\n\n\`\`\`json\n${JSON.stringify(params.tools, null, 2)}\n\`\`\`\n\n` - } - - if (params.fullHistory) { - markdown += `## Conversation History\n\n` - const [deletedStart, deletedEnd] = params.deletedRange || [-1, -1] - - for (let i = 0; i < params.fullHistory.length; i++) { - const message = params.fullHistory[i] - const isTruncated = i >= deletedStart && i <= deletedEnd - - markdown += `### [${message.role.toUpperCase()}]${isTruncated ? " [TRUNCATED]" : ""}\n` - - if (typeof message.content === "string") { - markdown += `${message.content}\n\n` - } else if (Array.isArray(message.content)) { - for (const block of message.content) { - if (block.type === "text") { - markdown += `**Text:** ${block.call_id ? `(\`call_id: ${block.call_id}\`)` : ""}\n${block.text}\n\n` - } else if (block.type === "thinking") { - markdown += `**Thinking:** ${block.call_id ? `(\`call_id: ${block.call_id}\`)` : ""}\n${block.thinking}\n\n` - } else if (block.type === "redacted_thinking") { - markdown += `**Thinking:** [Redacted] ${block.call_id ? `(\`call_id: ${block.call_id}\`)` : ""}\n\n` - } else if (block.type === "tool_use") { - markdown += `**Tool Use:** \`${block.name}\` (\`id: ${block.id}\`, \`call_id: ${block.call_id}\`)\n` - markdown += `\`\`\`json\n${JSON.stringify(block.input, null, 2)}\n\`\`\`\n\n` - } else if (block.type === "tool_result") { - markdown += `**Tool Result:** (\`${block.tool_use_id}\`)\n` - if (typeof block.content === "string") { - markdown += `${block.content}\n\n` - } else if (Array.isArray(block.content)) { - for (const contentBlock of block.content) { - if (contentBlock.type === "text") { - markdown += `${contentBlock.text}\n\n` - } else if (contentBlock.type === "image") { - markdown += `[Image: ${contentBlock.source?.type}]\n\n` - } - } - } - } else if (block.type === "image") { - markdown += `[Image: ${block.source?.type}]\n\n` - } - } - } - markdown += "---\n\n" - } - } - - await fs.writeFile(debugPath, markdown, "utf8") - } catch (error) { - Logger.error("Failed to write prompt metadata artifacts:", error) - } + return writePromptMetadataArtifacts(this.promptArtifactsContext, params) } private getApiRequestIdSafe(): string | undefined { @@ -1419,7 +1319,7 @@ export class Task { this.stateManager.getGlobalSettingsKey("useAutoCondense"), ) - await this.writePromptMetadataArtifacts({ + await writePromptMetadataArtifacts(this.promptArtifactsContext, { systemPrompt, providerInfo, tools: toolSnapshot.nativeTools, From 0a1b983da233a47c66834d642a6bce13fa1cf5f2 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 9 Aug 2026 15:03:05 +0000 Subject: [PATCH 3/8] FB-15a: extract buildApiRequestParams into TaskRequestBuilder.ts Co-Authored-By: Alexandros Salapatas --- src/core/task/TaskRequestBuilder.ts | 249 ++++++++++++++++++++++++++++ src/core/task/index.ts | 245 +++------------------------ 2 files changed, 276 insertions(+), 218 deletions(-) create mode 100644 src/core/task/TaskRequestBuilder.ts diff --git a/src/core/task/TaskRequestBuilder.ts b/src/core/task/TaskRequestBuilder.ts new file mode 100644 index 00000000..b2faaaea --- /dev/null +++ b/src/core/task/TaskRequestBuilder.ts @@ -0,0 +1,249 @@ +import type { ApiHandler, ApiProviderInfo } from "@core/api" +import { + getGlobalDiracRules, + getLocalDiracRules, + refreshDiracRulesToggles, +} from "@core/context/instructions/user-instructions/dirac-rules" +import { + getLocalAgentsRules, + getLocalCursorRules, + getLocalWindsurfRules, + refreshExternalRulesToggles, +} from "@core/context/instructions/user-instructions/external-rules" +import { formatResponse } from "@core/formatResponse" +import { ensureRulesDirectoryExists, ensureTaskDirectoryExists } from "@core/storage/disk" +import { createDefaultTextCondensationTemplateRegistry, TASK_HANDOFF_TEMPLATE_ID } from "@core/text-condensation/templates" +import { isUtilityTextCondensationAvailable } from "@core/text-condensation/UtilityTextCondensationAvailability" +import { isMultiRootEnabled } from "@core/workspace/multi-root-utils" +import { HostProvider } from "@hosts/host-provider" +import { featureFlagsService } from "@services/feature-flags" +import { DiracClient } from "@shared/dirac" +import { DEFAULT_LANGUAGE_SETTINGS, getLanguageKey, type LanguageDisplay } from "@shared/Languages" +import * as path from "path" +import { filterSkillsByProviderCapabilities } from "@/shared/skills" +import { getAvailableCores } from "@/utils/os" +import { detectBestShell } from "@/utils/shell-detection" +import type { ContextManager } from "../context/context-management/ContextManager" +import { RuleContextBuilder } from "../context/instructions/user-instructions/RuleContextBuilder" +import { getOrDiscoverSkills } from "../context/instructions/user-instructions/skills" +import type { DiracIgnoreController } from "../ignore/DiracIgnoreController" +import type { SystemPromptContext } from "../prompts/system-prompt" +import { getSystemPrompt } from "../prompts/system-prompt" +import type { StateManager } from "../storage/StateManager" +import type { WorkspaceRootManager } from "../workspace/WorkspaceRootManager" +import type { ApiConversationManager } from "./ApiConversationManager" +import type { MessageStateHandler } from "./message-state" +import type { TaskMessenger } from "./TaskMessenger" +import type { TaskState } from "./TaskState" +import type { ToolExecutor } from "./ToolExecutor" + +export interface TaskRequestBuilderContext { + taskId: string + cwd: string + terminalExecutionMode: "vscodeTerminal" | "backgroundExec" + api: ApiHandler + stateManager: StateManager + messageStateHandler: MessageStateHandler + taskMessenger: TaskMessenger + toolExecutor: ToolExecutor + contextManager: ContextManager + apiConversationManager: ApiConversationManager + diracIgnoreController: DiracIgnoreController + workspaceManager?: WorkspaceRootManager + taskState: TaskState + getCurrentProviderInfo: () => ApiProviderInfo + isParallelToolCallingEnabled: () => boolean +} + +export async function buildApiRequestParams( + ctx: TaskRequestBuilderContext, + params: { previousApiReqIndex: number; shouldCompact?: boolean }, +): Promise<{ + systemPrompt: string + toolSnapshot: Awaited> + contextManagementMetadata: Awaited> + providerInfo: ApiProviderInfo +}> { + const providerInfo = ctx.getCurrentProviderInfo() + const host = await HostProvider.env.getHostVersion({}) + const ide = host?.platform || "Unknown" + const isCliEnvironment = host.diracType === DiracClient.Cli + const browserSettings = ctx.stateManager.getGlobalSettingsKey("browserSettings") + const disableBrowserTool = browserSettings?.disableToolUse ?? false + const modelSupportsBrowserUse = providerInfo.model.info.supportsImages ?? false + + const supportsBrowserUse = modelSupportsBrowserUse && !disableBrowserTool + const preferredLanguageRaw = ctx.stateManager.getGlobalSettingsKey("preferredLanguage") + const preferredLanguage = getLanguageKey(preferredLanguageRaw as LanguageDisplay) + const preferredLanguageInstructions = + preferredLanguage && preferredLanguage !== DEFAULT_LANGUAGE_SETTINGS + ? `# Preferred Language\n\nSpeak in ${preferredLanguage}.` + : "" + + const { globalToggles, localToggles } = await refreshDiracRulesToggles(ctx.stateManager, ctx.cwd) + const { windsurfLocalToggles, cursorLocalToggles, agentsLocalToggles } = await refreshExternalRulesToggles( + ctx.stateManager, + ctx.cwd, + ) + + const evaluationContext = await new RuleContextBuilder().buildEvaluationContext({ + cwd: ctx.cwd, + messageStateHandler: ctx.messageStateHandler, + workspaceManager: ctx.workspaceManager, + }) + + const globalDiracRulesFilePath = await ensureRulesDirectoryExists() + const globalRules = await getGlobalDiracRules(globalDiracRulesFilePath, globalToggles, { evaluationContext }) + const globalDiracRulesFileInstructions = globalRules.instructions + const localRules = await getLocalDiracRules(ctx.cwd, localToggles, { evaluationContext }) + const localDiracRulesFileInstructions = localRules.instructions + const [localCursorRulesFileInstructions, localCursorRulesDirInstructions] = await getLocalCursorRules( + ctx.cwd, + cursorLocalToggles, + ) + const localWindsurfRulesFileInstructions = await getLocalWindsurfRules(ctx.cwd, windsurfLocalToggles) + const localAgentsRulesFileInstructions = await getLocalAgentsRules(ctx.cwd, agentsLocalToggles) + ctx.diracIgnoreController.yoloMode = !!ctx.stateManager.getGlobalSettingsKey("yoloModeToggled") + const isYolo = !!ctx.stateManager.getGlobalSettingsKey("yoloModeToggled") + const diracIgnoreContent = ctx.diracIgnoreController.diracIgnoreContent + let diracIgnoreInstructions: string | undefined + if (diracIgnoreContent && !isYolo) { + diracIgnoreInstructions = formatResponse.diracIgnoreInstructions(diracIgnoreContent) + } + + let workspaceRoots: Array<{ path: string; name: string; vcs?: string }> | undefined + const multiRootEnabled = isMultiRootEnabled(ctx.stateManager) + if (multiRootEnabled && ctx.workspaceManager) { + workspaceRoots = ctx.workspaceManager.getRoots().map((root) => ({ + path: root.path, + name: root.name || path.basename(root.path), + vcs: root.vcs as string | undefined, + })) + } + + const resolvedSkills = await getOrDiscoverSkills(ctx.cwd, ctx.taskState) + const providerSkills = filterSkillsByProviderCapabilities(resolvedSkills, { + native_web_search: providerInfo.supportsNativeWebSearch === true, + }) + const globalSkillsToggles = ctx.stateManager.getGlobalSettingsKey("globalSkillsToggles") ?? {} + const localSkillsToggles = ctx.stateManager.getWorkspaceStateKey("localSkillsToggles") ?? {} + const availableSkills = providerSkills.filter((skill) => { + if (ctx.stateManager.getGlobalSettingsKey("yoloModeToggled") && skill.interactiveOnly) return false + if (skill.source === "builtin") return true + const toggles = skill.source === "global" ? globalSkillsToggles : localSkillsToggles + return toggles[skill.path] !== false + }) + ctx.taskState.availableSkills = availableSkills + + const openTabPaths = (await HostProvider.window.getOpenTabs({})).paths || [] + const visibleTabPaths = (await HostProvider.window.getVisibleTabs({})).paths || [] + const cap = 50 + const editorTabs = { + open: openTabPaths.slice(0, cap), + visible: visibleTabPaths.slice(0, cap), + } + const shellInfo = detectBestShell() + const taskHandoffCondensationAvailable = isUtilityTextCondensationAvailable( + { + utilityModelEnabled: ctx.stateManager.getGlobalSettingsKey("utilityModelEnabled"), + utilityModelUseCondense: ctx.stateManager.getGlobalSettingsKey("utilityModelUseCondense"), + utilityModelUseNewTask: ctx.stateManager.getGlobalSettingsKey("utilityModelUseNewTask"), + utilityModelSelection: ctx.stateManager.getGlobalSettingsKey("utilityModelSelection"), + }, + TASK_HANDOFF_TEMPLATE_ID, + createDefaultTextCondensationTemplateRegistry(), + ) + + const promptContext: SystemPromptContext = { + cwd: ctx.cwd, + ide, + providerInfo, + editorTabs, + supportsBrowserUse, + taskHandoffCondensationAvailable, + skills: availableSkills, + globalDiracRulesFileInstructions, + localDiracRulesFileInstructions, + localCursorRulesFileInstructions, + localCursorRulesDirInstructions, + localWindsurfRulesFileInstructions, + localAgentsRulesFileInstructions, + diracIgnoreInstructions, + preferredLanguageInstructions, + browserSettings: ctx.stateManager.getGlobalSettingsKey("browserSettings"), + yoloModeToggled: ctx.stateManager.getGlobalSettingsKey("yoloModeToggled"), + subagentsEnabled: ctx.stateManager.getGlobalSettingsKey("subagentsEnabled"), + diracWebToolsEnabled: + ctx.stateManager.getGlobalSettingsKey("diracWebToolsEnabled") && featureFlagsService.getWebtoolsEnabled(), + isMultiRootEnabled: multiRootEnabled, + workspaceRoots, + isSubagentRun: false, + isCliEnvironment, + enableParallelToolCalling: ctx.isParallelToolCallingEnabled(), + terminalExecutionMode: ctx.terminalExecutionMode, + activeShellType: shellInfo.type, + activeShellPath: shellInfo.path, + activeShellIsPosix: shellInfo.isPosix, + availableCores: getAvailableCores(), + shouldCompact: params.shouldCompact, + } + + const activatedConditionalRules = [...globalRules.activatedConditionalRules, ...localRules.activatedConditionalRules] + if (activatedConditionalRules.length > 0) { + await ctx.taskMessenger.upsertText(JSON.stringify({ rules: activatedConditionalRules })) + } + + const ruleLoadErrors = [...(globalRules.errors ?? []), ...(localRules.errors ?? [])] + if (ruleLoadErrors.length > 0) { + await ctx.taskMessenger.upsertText(JSON.stringify({ ruleLoadErrors })) + } + + const toolSnapshot = await ctx.toolExecutor.getSnapshotForRequest(promptContext) + const { systemPrompt } = await getSystemPrompt(promptContext, toolSnapshot) + ctx.toolExecutor.activateSnapshot(toolSnapshot) + ctx.taskState.useNativeToolCalls = toolSnapshot.nativeTools.length > 0 + const contextManagementMetadata = await ctx.contextManager.getNewContextMessagesAndMetadata( + ctx.messageStateHandler.getApiConversationHistory(), + ctx.messageStateHandler.getDiracMessages(), + ctx.api, + ctx.taskState.conversationHistoryDeletedRange, + params.previousApiReqIndex, + await ensureTaskDirectoryExists(ctx.taskId), + ctx.stateManager.getGlobalSettingsKey("useAutoCondense"), + ) + + if (contextManagementMetadata.updatedConversationHistoryDeletedRange) { + const previousConversationHistoryDeletedRange = ctx.taskState.conversationHistoryDeletedRange + const conversationHistoryDeletedRange = contextManagementMetadata.conversationHistoryDeletedRange + if (!conversationHistoryDeletedRange) { + throw new Error("Context management reported a truncation update without a deleted range.") + } + ctx.taskState.conversationHistoryDeletedRange = conversationHistoryDeletedRange + await ctx.apiConversationManager.scheduleProviderConversationCompaction( + previousConversationHistoryDeletedRange, + conversationHistoryDeletedRange, + ) + await ctx.messageStateHandler.saveDiracMessagesAndUpdateHistory() + } + + const useAutoCondense = ctx.stateManager.getGlobalSettingsKey("useAutoCondense") + if (!useAutoCondense) { + const lastMessage = + contextManagementMetadata.truncatedConversationHistory[ + contextManagementMetadata.truncatedConversationHistory.length - 1 + ] + if (lastMessage && lastMessage.role === "user") { + const notice = formatResponse.contextTruncationNotice() + if (typeof lastMessage.content === "string") { + lastMessage.content += `\n\n${notice}` + } else if (Array.isArray(lastMessage.content)) { + lastMessage.content.push({ + type: "text", + text: notice, + }) + } + } + } + + return { systemPrompt, toolSnapshot, contextManagementMetadata, providerInfo } +} diff --git a/src/core/task/index.ts b/src/core/task/index.ts index be521948..1cfed395 100644 --- a/src/core/task/index.ts +++ b/src/core/task/index.ts @@ -5,26 +5,12 @@ import { ContextManager } from "@core/context/context-management/ContextManager" import { EnvironmentContextTracker } from "@core/context/context-tracking/EnvironmentContextTracker" import { FileContextTracker } from "@core/context/context-tracking/FileContextTracker" import { ModelContextTracker } from "@core/context/context-tracking/ModelContextTracker" -import { - getGlobalDiracRules, - getLocalDiracRules, - refreshDiracRulesToggles, -} from "@core/context/instructions/user-instructions/dirac-rules" -import { - getLocalAgentsRules, - getLocalCursorRules, - getLocalWindsurfRules, - refreshExternalRulesToggles, -} from "@core/context/instructions/user-instructions/external-rules" import { formatResponse } from "@core/formatResponse" import { DiracIgnoreController } from "@core/ignore/DiracIgnoreController" import { recordSuccessfulModelProviderPreset } from "@core/models/modelProviderPresets" import { CommandPermissionController } from "@core/permissions" -import type { SystemPromptContext } from "@core/prompts/system-prompt" -import { getSystemPrompt } from "@core/prompts/system-prompt" import type { SlashCommandDirectAction } from "@core/slash-commands" -import { ensureRulesDirectoryExists, ensureTaskDirectoryExists } from "@core/storage/disk" -import { createDefaultTextCondensationTemplateRegistry, TASK_HANDOFF_TEMPLATE_ID } from "@core/text-condensation/templates" +import { createDefaultTextCondensationTemplateRegistry } from "@core/text-condensation/templates" import { isUtilityTextCondensationAvailable } from "@core/text-condensation/UtilityTextCondensationAvailability" import { getConfiguredUtilityModelSelection } from "@core/utility-model/UtilityModelSelection" import { isMultiRootEnabled } from "@core/workspace/multi-root-utils" @@ -48,15 +34,14 @@ import { ITerminalManager } from "@integrations/terminal/types" import { BrowserSession } from "@services/browser/BrowserSession" import { UrlContentFetcher } from "@services/browser/UrlContentFetcher" import { DiracError, DiracErrorType, ErrorService } from "@services/error" -import { featureFlagsService } from "@services/feature-flags" + import { telemetryService } from "@services/telemetry" import { ApiConfiguration } from "@shared/api" import { findLastIndex } from "@shared/array" -import { DiracClient } from "@shared/dirac" import { getExtensionSourceDir } from "@shared/dirac/constants" import { CardStatus, DiracApiReqCancelReason, DiracMessageContent, DiracMessageType, TaskStatus } from "@shared/ExtensionMessage" import { HistoryItem } from "@shared/HistoryItem" -import { DEFAULT_LANGUAGE_SETTINGS, getLanguageKey, LanguageDisplay } from "@shared/Languages" + import { DiracContent, DiracStorageMessage, @@ -75,14 +60,10 @@ import { DiracAskResponse } from "@shared/WebviewMessage" import { isLocalModel, isParallelToolCallingEnabled } from "@utils/model-utils" import Mutex from "p-mutex" import pWaitFor from "p-wait-for" -import * as path from "path" import { ulid } from "ulid" import { getErrorMessage } from "@/shared/errors" -import { filterSkillsByProviderCapabilities, SkillMetadata } from "@/shared/skills" -import { getAvailableCores } from "@/utils/os" -import { detectBestShell } from "@/utils/shell-detection" -import { RuleContextBuilder } from "../context/instructions/user-instructions/RuleContextBuilder" -import { getOrDiscoverSkills } from "../context/instructions/user-instructions/skills" +import { type SkillMetadata } from "@/shared/skills" + import { Controller } from "../controller" import { StateManager } from "../storage/StateManager" import { ApiConversationManager } from "./ApiConversationManager" @@ -101,6 +82,7 @@ import { StreamResponseHandler } from "./StreamResponseHandler" import { type SteeringClaim } from "./steering" import { TaskMessenger } from "./TaskMessenger" import { type TaskPromptArtifactsContext, writePromptMetadataArtifacts } from "./TaskPromptArtifacts" +import { buildApiRequestParams, type TaskRequestBuilderContext } from "./TaskRequestBuilder" import { TaskState } from "./TaskState" import { appendQueuedSteeringToNextApiRequest, @@ -190,6 +172,26 @@ export class Task { } } + private get requestBuilderContext(): TaskRequestBuilderContext { + return { + taskId: this.taskId, + cwd: this.cwd, + terminalExecutionMode: this.terminalExecutionMode, + api: this.api, + stateManager: this.stateManager, + messageStateHandler: this.messageStateHandler, + taskMessenger: this.taskMessenger, + toolExecutor: this.toolExecutor, + contextManager: this.contextManager, + apiConversationManager: this.apiConversationManager, + diracIgnoreController: this.diracIgnoreController, + workspaceManager: this.workspaceManager, + taskState: this.taskState, + getCurrentProviderInfo: () => this.getCurrentProviderInfo(), + isParallelToolCallingEnabled: () => this.isParallelToolCallingEnabled(), + } + } + public canAcceptSteeringMessage(): boolean { return canAcceptSteeringMessage(this.steeringContext) } @@ -1168,200 +1170,7 @@ export class Task { * for an API request. Extracted from attemptApiRequest for single-responsibility. */ private async buildApiRequestParams(params: { previousApiReqIndex: number; shouldCompact?: boolean }) { - const providerInfo = this.getCurrentProviderInfo() - const host = await HostProvider.env.getHostVersion({}) - const ide = host?.platform || "Unknown" - const isCliEnvironment = host.diracType === DiracClient.Cli - const browserSettings = this.stateManager.getGlobalSettingsKey("browserSettings") - const disableBrowserTool = browserSettings.disableToolUse ?? false - // dirac browser tool uses image recognition for navigation (requires model image support). - const modelSupportsBrowserUse = providerInfo.model.info.supportsImages ?? false - - const supportsBrowserUse = modelSupportsBrowserUse && !disableBrowserTool // only enable browser use if the model supports it and the user hasn't disabled it - const preferredLanguageRaw = this.stateManager.getGlobalSettingsKey("preferredLanguage") - const preferredLanguage = getLanguageKey(preferredLanguageRaw as LanguageDisplay) - const preferredLanguageInstructions = - preferredLanguage && preferredLanguage !== DEFAULT_LANGUAGE_SETTINGS - ? `# Preferred Language\n\nSpeak in ${preferredLanguage}.` - : "" - - const { globalToggles, localToggles } = await refreshDiracRulesToggles(this.stateManager, this.cwd) - const { windsurfLocalToggles, cursorLocalToggles, agentsLocalToggles } = await refreshExternalRulesToggles( - this.stateManager, - this.cwd, - ) - - const evaluationContext = await new RuleContextBuilder().buildEvaluationContext({ - cwd: this.cwd, - messageStateHandler: this.messageStateHandler, - workspaceManager: this.workspaceManager, - }) - - const globalDiracRulesFilePath = await ensureRulesDirectoryExists() - const globalRules = await getGlobalDiracRules(globalDiracRulesFilePath, globalToggles, { evaluationContext }) - const globalDiracRulesFileInstructions = globalRules.instructions - const localRules = await getLocalDiracRules(this.cwd, localToggles, { evaluationContext }) - const localDiracRulesFileInstructions = localRules.instructions - const [localCursorRulesFileInstructions, localCursorRulesDirInstructions] = await getLocalCursorRules( - this.cwd, - cursorLocalToggles, - ) - const localWindsurfRulesFileInstructions = await getLocalWindsurfRules(this.cwd, windsurfLocalToggles) - const localAgentsRulesFileInstructions = await getLocalAgentsRules(this.cwd, agentsLocalToggles) - this.diracIgnoreController.yoloMode = !!this.stateManager.getGlobalSettingsKey("yoloModeToggled") - const isYolo = !!this.stateManager.getGlobalSettingsKey("yoloModeToggled") - const diracIgnoreContent = this.diracIgnoreController.diracIgnoreContent - let diracIgnoreInstructions: string | undefined - if (diracIgnoreContent && !isYolo) { - diracIgnoreInstructions = formatResponse.diracIgnoreInstructions(diracIgnoreContent) - } - // Prepare multi-root workspace information if enabled - let workspaceRoots: Array<{ path: string; name: string; vcs?: string }> | undefined - const multiRootEnabled = isMultiRootEnabled(this.stateManager) - if (multiRootEnabled && this.workspaceManager) { - workspaceRoots = this.workspaceManager.getRoots().map((root) => ({ - path: root.path, - name: root.name || path.basename(root.path), // Fallback to basename if name is undefined - vcs: root.vcs as string | undefined, // Cast VcsType to string - })) - } - // Discover and filter available skills - const resolvedSkills = await getOrDiscoverSkills(this.cwd, this.taskState) - const providerSkills = filterSkillsByProviderCapabilities(resolvedSkills, { - native_web_search: providerInfo.supportsNativeWebSearch === true, - }) - // Filter skills by toggle state (enabled by default) - const globalSkillsToggles = this.stateManager.getGlobalSettingsKey("globalSkillsToggles") ?? {} - const localSkillsToggles = this.stateManager.getWorkspaceStateKey("localSkillsToggles") ?? {} - const availableSkills = providerSkills.filter((skill) => { - if (this.stateManager.getGlobalSettingsKey("yoloModeToggled") && skill.interactiveOnly) return false - if (skill.source === "builtin") return true - const toggles = skill.source === "global" ? globalSkillsToggles : localSkillsToggles - return toggles[skill.path] !== false - }) - this.taskState.availableSkills = availableSkills - // Snapshot editor tabs so prompt tools can decide whether to include - // filetype-specific instructions (e.g. notebooks) without adding bespoke flags. - const openTabPaths = (await HostProvider.window.getOpenTabs({})).paths || [] - const visibleTabPaths = (await HostProvider.window.getVisibleTabs({})).paths || [] - const cap = 50 - const editorTabs = { - open: openTabPaths.slice(0, cap), - visible: visibleTabPaths.slice(0, cap), - } - const shellInfo = detectBestShell() - const taskHandoffCondensationAvailable = isUtilityTextCondensationAvailable( - { - utilityModelEnabled: this.stateManager.getGlobalSettingsKey("utilityModelEnabled"), - utilityModelUseCondense: this.stateManager.getGlobalSettingsKey("utilityModelUseCondense"), - utilityModelUseNewTask: this.stateManager.getGlobalSettingsKey("utilityModelUseNewTask"), - utilityModelSelection: this.stateManager.getGlobalSettingsKey("utilityModelSelection"), - }, - TASK_HANDOFF_TEMPLATE_ID, - createDefaultTextCondensationTemplateRegistry(), - ) - const promptContext: SystemPromptContext = { - cwd: this.cwd, - ide, - providerInfo, - editorTabs, - supportsBrowserUse, - taskHandoffCondensationAvailable, - skills: availableSkills, - globalDiracRulesFileInstructions, - localDiracRulesFileInstructions, - localCursorRulesFileInstructions, - localCursorRulesDirInstructions, - localWindsurfRulesFileInstructions, - localAgentsRulesFileInstructions, - diracIgnoreInstructions, - preferredLanguageInstructions, - browserSettings: this.stateManager.getGlobalSettingsKey("browserSettings"), - yoloModeToggled: this.stateManager.getGlobalSettingsKey("yoloModeToggled"), - subagentsEnabled: this.stateManager.getGlobalSettingsKey("subagentsEnabled"), - utilityModelConfigured: - getConfiguredUtilityModelSelection(this.stateManager.getGlobalSettingsKey("utilityModelSelection")) !== undefined, - diracWebToolsEnabled: - this.stateManager.getGlobalSettingsKey("diracWebToolsEnabled") && featureFlagsService.getWebtoolsEnabled(), - isMultiRootEnabled: multiRootEnabled, - workspaceRoots, - isSubagentRun: false, - isCliEnvironment, - enableParallelToolCalling: this.isParallelToolCallingEnabled(), - terminalExecutionMode: this.terminalExecutionMode, - activeShellType: shellInfo.type, - activeShellPath: shellInfo.path, - activeShellIsPosix: shellInfo.isPosix, - availableCores: getAvailableCores(), - shouldCompact: params.shouldCompact, - } - // Notify user if any conditional rules were applied for this request - const activatedConditionalRules = [...globalRules.activatedConditionalRules, ...localRules.activatedConditionalRules] - if (activatedConditionalRules.length > 0) { - await this.taskMessenger.upsertText(JSON.stringify({ rules: activatedConditionalRules })) - } - // Surface rule-file load failures so silently-dropped rules are visible to the user. - const ruleLoadErrors = [...(globalRules.errors ?? []), ...(localRules.errors ?? [])] - if (ruleLoadErrors.length > 0) { - await this.taskMessenger.upsertText(JSON.stringify({ ruleLoadErrors })) - } - const toolSnapshot = await this.toolExecutor.getSnapshotForRequest(promptContext) - const { systemPrompt } = await getSystemPrompt(promptContext, toolSnapshot) - this.toolExecutor.activateSnapshot(toolSnapshot) - this.taskState.useNativeToolCalls = toolSnapshot.nativeTools.length > 0 - const contextManagementMetadata = await this.contextManager.getNewContextMessagesAndMetadata( - this.messageStateHandler.getApiConversationHistory(), - this.messageStateHandler.getDiracMessages(), - this.api, - this.taskState.conversationHistoryDeletedRange, - params.previousApiReqIndex, - await ensureTaskDirectoryExists(this.taskId), - this.stateManager.getGlobalSettingsKey("useAutoCondense"), - ) - - await writePromptMetadataArtifacts(this.promptArtifactsContext, { - systemPrompt, - providerInfo, - tools: toolSnapshot.nativeTools, - fullHistory: this.messageStateHandler.getApiConversationHistory(), - deletedRange: this.taskState.conversationHistoryDeletedRange, - }) - - if (contextManagementMetadata.updatedConversationHistoryDeletedRange) { - const previousConversationHistoryDeletedRange = this.taskState.conversationHistoryDeletedRange - const conversationHistoryDeletedRange = contextManagementMetadata.conversationHistoryDeletedRange - if (!conversationHistoryDeletedRange) { - throw new Error("Context management reported a truncation update without a deleted range.") - } - this.taskState.conversationHistoryDeletedRange = conversationHistoryDeletedRange - await this.apiConversationManager.scheduleProviderConversationCompaction( - previousConversationHistoryDeletedRange, - conversationHistoryDeletedRange, - ) - await this.messageStateHandler.saveDiracMessagesAndUpdateHistory() - } - - // If we're not using auto-condense, we should explicitly notify the model that history was truncated - const useAutoCondense = this.stateManager.getGlobalSettingsKey("useAutoCondense") - if (!useAutoCondense) { - const lastMessage = - contextManagementMetadata.truncatedConversationHistory[ - contextManagementMetadata.truncatedConversationHistory.length - 1 - ] - if (lastMessage && lastMessage.role === "user") { - const notice = formatResponse.contextTruncationNotice() - if (typeof lastMessage.content === "string") { - lastMessage.content += `\n\n${notice}` - } else if (Array.isArray(lastMessage.content)) { - lastMessage.content.push({ - type: "text", - text: notice, - }) - } - } - } - - return { systemPrompt, toolSnapshot, contextManagementMetadata, providerInfo } + return buildApiRequestParams(this.requestBuilderContext, params) } private async handleApiRequestError(params: { From a891e9bf3d1b99b4232c95e50aee877878d297be Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 9 Aug 2026 15:08:45 +0000 Subject: [PATCH 4/8] FB-15a: extract handleApiRequestError and processStreamResult into TaskRequestOutcome.ts Co-Authored-By: Alexandros Salapatas --- src/core/task/TaskRequestOutcome.ts | 334 ++++++++++++++++++++++++++++ src/core/task/index.ts | 289 +++--------------------- 2 files changed, 360 insertions(+), 263 deletions(-) create mode 100644 src/core/task/TaskRequestOutcome.ts diff --git a/src/core/task/TaskRequestOutcome.ts b/src/core/task/TaskRequestOutcome.ts new file mode 100644 index 00000000..9aa54d81 --- /dev/null +++ b/src/core/task/TaskRequestOutcome.ts @@ -0,0 +1,334 @@ +import { setTimeout as setTimeoutPromise } from "node:timers/promises" +import type { ApiHandler } from "@core/api" +import type { ApiStream } from "@core/api/transform/stream" +import { formatResponse } from "@core/formatResponse" +import type { ICheckpointManager } from "@integrations/checkpoints/types" +import { processFilesIntoText } from "@integrations/misc/extract-text" +import { DiracError, DiracErrorType } from "@services/error" +import { findLastIndex } from "@shared/array" +import { CardStatus, DiracMessageType, TaskStatus } from "@shared/ExtensionMessage" +import { DiracAskResponse } from "@shared/WebviewMessage" +import { DiracContent, DiracTextContentBlock } from "@shared/messages/content" +import type { DiracMessageModelInfo } from "@shared/messages/metrics" +import { isMutatingTool } from "@shared/tools" +import pWaitFor from "p-wait-for" +import type { MessageStateHandler } from "./message-state" +import type { StreamingMetricsManager } from "./StreamingMetricsManager" +import type { TaskMessenger } from "./TaskMessenger" +import type { TaskState } from "./TaskState" +import { ToolSkippedByUserMessage } from "./tools/types/ToolSkippedByUserMessage" +import { updateApiReqMsg } from "./utils" + +export interface TaskRequestOutcomeContext { + taskState: TaskState + messageStateHandler: MessageStateHandler + taskMessenger: TaskMessenger + api: ApiHandler + taskId: string + checkpointManager?: ICheckpointManager + postStateToWebview: () => Promise + abortTask: () => Promise + handleContextWindowExceededError: () => Promise + reinitExistingTaskFromId: () => Promise + attemptApiRequest: (previousApiReqIndex: number, lastApiReqIndex: number, shouldCompact?: boolean) => ApiStream + recursivelyMakeDiracRequests: (userContent: DiracContent[], includeFileDetails?: boolean) => Promise + handleEmptyAssistantResponse: (params: { + modelInfo: DiracMessageModelInfo + taskMetrics: { + inputTokens: number + outputTokens: number + cacheWriteTokens: number + cacheReadTokens: number + totalCost?: number + } + providerId: string + model: { id: string } + }) => Promise +} + +export async function handleApiRequestError( + ctx: TaskRequestOutcomeContext, + params: { + error: unknown + previousApiReqIndex: number + lastApiReqIndex: number + shouldCompact?: boolean + model: { id: string; info: { contextWindow?: number } } + providerId: string + metricsManager: StreamingMetricsManager + }, +): Promise { + const { error, model, providerId } = params + const diracError = DiracError.transform(error, model.id, providerId) + + if (diracError.isErrorType(DiracErrorType.ContextWindowExceeded)) { + await ctx.handleContextWindowExceededError() + const truncatedConversationHistory = ctx.messageStateHandler.getDiracMessages() + if (truncatedConversationHistory.length > 3) { + diracError.message = "Context window exceeded. Click retry to truncate the conversation and try again." + } + } + + const streamingFailedMessage = diracError.serialize() + + const lastApiReqStartedIndex = findLastIndex( + ctx.messageStateHandler.getDiracMessages(), + (m) => m.content.type === DiracMessageType.API_STATUS, + ) + if (lastApiReqStartedIndex !== -1) { + const diracMessages = ctx.messageStateHandler.getDiracMessages() + const msg = diracMessages[lastApiReqStartedIndex] + if (msg.content.type === DiracMessageType.API_STATUS) { + const currentApiReqInfo = { ...msg.content.status } + delete currentApiReqInfo.retryStatus + + await ctx.messageStateHandler.updateDiracMessage(lastApiReqStartedIndex, { + content: { + type: DiracMessageType.API_STATUS, + status: { + ...currentApiReqInfo, + streamingFailedMessage, + }, + }, + }) + } + } + + const isAuthError = diracError.isErrorType(DiracErrorType.Auth) + const isPaymentError = diracError.isErrorType(DiracErrorType.Payment) + + let response: DiracAskResponse + if (!isAuthError && !isPaymentError && ctx.taskState.apiErrorRetryAttempts < 3) { + ctx.taskState.apiErrorRetryAttempts++ + const delay = 2000 * 2 ** (ctx.taskState.apiErrorRetryAttempts - 1) + + await updateApiReqMsg({ + messageStateHandler: ctx.messageStateHandler, + lastApiReqIndex: lastApiReqStartedIndex, + inputTokens: 0, + reasoningTokens: 0, + outputTokens: 0, + cacheWriteTokens: 0, + cacheReadTokens: 0, + totalCost: undefined, + api: ctx.api, + cancelReason: "streaming_failed", + streamingFailedMessage, + }) + await ctx.messageStateHandler.saveDiracMessagesAndUpdateHistory() + await ctx.postStateToWebview() + + response = DiracAskResponse.APPROVE + const autoRetryCard = await ctx.taskMessenger.createCard({ + status: CardStatus.PENDING, + header: "API Error (Retrying)", + body: `API Error (attempt ${ctx.taskState.apiErrorRetryAttempts}/3). Retrying in ${delay / 1000}s...`, + }) + + const autoRetryApiReqIndex = findLastIndex( + ctx.messageStateHandler.getDiracMessages(), + (m) => m.content.type === DiracMessageType.API_STATUS, + ) + if (autoRetryApiReqIndex !== -1) { + const diracMessages = ctx.messageStateHandler.getDiracMessages() + const msg = diracMessages[autoRetryApiReqIndex] + if (msg.content.type === DiracMessageType.API_STATUS) { + const currentApiReqInfo = { ...msg.content.status } + delete currentApiReqInfo.streamingFailedMessage + await ctx.messageStateHandler.updateDiracMessage(autoRetryApiReqIndex, { + content: { + type: DiracMessageType.API_STATUS, + status: currentApiReqInfo, + }, + }) + } + } + + const deadline = Date.now() + delay + while (Date.now() < deadline && !ctx.taskState.abort) { + await setTimeoutPromise(Math.min(200, deadline - Date.now())) + } + // If the user aborted during the retry delay, stop retrying + if (ctx.taskState.abort) { + await autoRetryCard.update({ + header: "API Error (Cancelled)", + body: `API Error (attempt ${ctx.taskState.apiErrorRetryAttempts}/3). Cancelled.`, + }) + await autoRetryCard.finalize(CardStatus.CANCELLED) + throw new Error("Task instance aborted") + } + await autoRetryCard.update({ body: `API Error (attempt ${ctx.taskState.apiErrorRetryAttempts}/3). Retrying...` }) + await autoRetryCard.finalize(CardStatus.ERROR) + } else { + if (!isAuthError && !isPaymentError) { + await ctx.taskMessenger.createCard({ + status: CardStatus.ERROR, + header: "API Error (Retries Exhausted)", + body: `The API request failed after 3 attempts. ${diracError.toDisplayMessage()}`, + }) + } + if (isPaymentError) { + await ctx.taskMessenger.createCard({ + status: CardStatus.ERROR, + header: "API Error (Payment Required)", + body: diracError.toDisplayMessage(), + }) + } + if (isAuthError) { + await ctx.taskMessenger.createCard({ + status: CardStatus.ERROR, + header: "API Error (Authentication)", + body: diracError.toDisplayMessage(), + }) + } + ctx.taskState.status = TaskStatus.AWAITING_USER_INPUT + + const cardHandle = await ctx.taskMessenger.createCard({ + requireApproval: true, + header: "API Request Failed", + body: diracError.toDisplayMessage(), + actions: [ + { label: "Retry", value: DiracAskResponse.APPROVE, primary: true }, + { label: "Cancel", value: DiracAskResponse.REJECT }, + ], + }) + try { + const askResult = await cardHandle.waitForInteraction() + response = askResult.response + } catch (error) { + if (error instanceof ToolSkippedByUserMessage) { + await cardHandle.finalize(CardStatus.SKIPPED) + ctx.taskState.pendingUserMessage = error.userMessage + ctx.taskState.pendingUserImages = error.userImages + ctx.taskState.pendingUserFiles = error.userFiles + response = DiracAskResponse.APPROVE + } else { + throw error + } + } + if (response === DiracAskResponse.APPROVE) { + ctx.taskState.apiErrorRetryAttempts = 0 + } + } + + if (response !== DiracAskResponse.APPROVE) { + await ctx.abortTask() + await ctx.reinitExistingTaskFromId() + return false + } + + const manualRetryApiReqIndex = findLastIndex( + ctx.messageStateHandler.getDiracMessages(), + (m) => m.content.type === DiracMessageType.API_STATUS, + ) + if (manualRetryApiReqIndex !== -1) { + const diracMessages = ctx.messageStateHandler.getDiracMessages() + const msg = diracMessages[manualRetryApiReqIndex] + if (msg.content.type === DiracMessageType.API_STATUS) { + const currentApiReqInfo = { ...msg.content.status } + delete currentApiReqInfo.streamingFailedMessage + await ctx.messageStateHandler.updateDiracMessage(manualRetryApiReqIndex, { + content: { + type: DiracMessageType.API_STATUS, + status: currentApiReqInfo, + }, + }) + } + } + + await ctx.taskMessenger.upsertText("Retrying API request...") + + return true +} + +export async function processStreamResult( + ctx: TaskRequestOutcomeContext, + params: { + assistantHasContent: boolean + stopReason?: string + userContent: DiracContent[] + metricsManager: StreamingMetricsManager + modelInfo: DiracMessageModelInfo + providerId: string + model: { id: string } + }, +): Promise { + if (params.assistantHasContent) { + ctx.taskState.askResponse = undefined + ctx.taskState.askResponseText = undefined + ctx.taskState.askResponseImages = undefined + ctx.taskState.askResponseFiles = undefined + ctx.taskState.status = TaskStatus.AWAITING_USER_INPUT + + await pWaitFor(() => ctx.taskState.userMessageContentReady) + const hasMutatingTools = ctx.taskState.assistantMessageContent.some( + (block) => block.type === "tool_use" && isMutatingTool(block.name), + ) + if (hasMutatingTools) { + await ctx.checkpointManager?.saveCheckpoint() + } + + const didToolUse = ctx.taskState.assistantMessageContent.some((block) => block.type === "tool_use") + if (ctx.taskState.didAttemptCompletion) { + ctx.taskState.completionCommitted = false + ctx.taskState.status = TaskStatus.COMPLETED + await ctx.postStateToWebview() + return true + } + const hitTokenLimit = + params.stopReason === "MAX_TOKENS" || params.stopReason === "max_tokens" || params.stopReason === "length" + + if (!didToolUse) { + ctx.taskState.userMessageContent.push({ + type: "text", + text: hitTokenLimit + ? "You reached the output token limit. Continue from where you stopped; restart an interrupted tool call, or call respond with operation 'complete' if finished." + : formatResponse.noToolsUsed(ctx.taskState.useNativeToolCalls), + } as DiracTextContentBlock) + ctx.taskState.consecutiveMistakeCount++ + } + + ctx.taskState.apiErrorRetryAttempts = 0 + ctx.taskState.emptyResponseRetryAttempts = 0 + + if ( + ctx.taskState.pendingUserMessage !== undefined || + (ctx.taskState.pendingUserImages?.length ?? 0) > 0 || + (ctx.taskState.pendingUserFiles?.length ?? 0) > 0 + ) { + if (ctx.taskState.pendingUserMessage) { + ctx.taskState.userMessageContent.push({ + type: "text", + text: `\n${ctx.taskState.pendingUserMessage}\n`, + isUserInput: true, + } as DiracTextContentBlock) + } + if (ctx.taskState.pendingUserImages?.length) { + ctx.taskState.userMessageContent.push(...formatResponse.imageBlocks(ctx.taskState.pendingUserImages)) + } + if (ctx.taskState.pendingUserFiles?.length) { + const fileContent = await processFilesIntoText(ctx.taskState.pendingUserFiles) + if (fileContent) { + ctx.taskState.userMessageContent.push({ type: "text", text: fileContent } as DiracTextContentBlock) + } + } + ctx.taskState.pendingUserMessage = undefined + ctx.taskState.pendingUserImages = undefined + ctx.taskState.pendingUserFiles = undefined + } + + return await ctx.recursivelyMakeDiracRequests(ctx.taskState.userMessageContent) + } + const taskMetrics = params.metricsManager.getMetrics() + const shouldRetry = await ctx.handleEmptyAssistantResponse({ + modelInfo: params.modelInfo, + taskMetrics, + providerId: params.providerId, + model: params.model, + }) + if (shouldRetry === false) { + ctx.taskState.consecutiveMistakeCount = 0 + return await ctx.recursivelyMakeDiracRequests(params.userContent) + } + return true +} diff --git a/src/core/task/index.ts b/src/core/task/index.ts index 1cfed395..91d0ec66 100644 --- a/src/core/task/index.ts +++ b/src/core/task/index.ts @@ -1,4 +1,3 @@ -import { setTimeout as setTimeoutPromise } from "node:timers/promises" import { ApiHandler, ApiProviderInfo, buildApiHandler } from "@core/api" import { ApiStream } from "@core/api/transform/stream" import { ContextManager } from "@core/context/context-management/ContextManager" @@ -33,7 +32,7 @@ import { import { ITerminalManager } from "@integrations/terminal/types" import { BrowserSession } from "@services/browser/BrowserSession" import { UrlContentFetcher } from "@services/browser/UrlContentFetcher" -import { DiracError, DiracErrorType, ErrorService } from "@services/error" +import { ErrorService } from "@services/error" import { telemetryService } from "@services/telemetry" import { ApiConfiguration } from "@shared/api" @@ -55,7 +54,7 @@ import { ShowMessageType } from "@shared/proto/index.host" import { Logger } from "@shared/services/Logger" import { Session } from "@shared/services/Session" import { type Mode } from "@shared/storage/types" -import { isMutatingTool } from "@shared/tools" + import { DiracAskResponse } from "@shared/WebviewMessage" import { isLocalModel, isParallelToolCallingEnabled } from "@utils/model-utils" import Mutex from "p-mutex" @@ -83,6 +82,7 @@ import { type SteeringClaim } from "./steering" import { TaskMessenger } from "./TaskMessenger" import { type TaskPromptArtifactsContext, writePromptMetadataArtifacts } from "./TaskPromptArtifacts" import { buildApiRequestParams, type TaskRequestBuilderContext } from "./TaskRequestBuilder" +import { handleApiRequestError, processStreamResult, type TaskRequestOutcomeContext } from "./TaskRequestOutcome" import { TaskState } from "./TaskState" import { appendQueuedSteeringToNextApiRequest, @@ -101,7 +101,7 @@ import { ToolExecutor } from "./ToolExecutor" import { DiracContext } from "./tools/context/DiracContext" import type { ToolSnapshotDirtyReason } from "./tools/runtime/ToolSnapshot" import { ToolSkippedByUserMessage } from "./tools/types/ToolSkippedByUserMessage" -import { extractProviderDomainFromUrl, updateApiReqMsg } from "./utils" +import { extractProviderDomainFromUrl } from "./utils" export type ToolResponse = DiracToolResponseContent @@ -192,6 +192,26 @@ export class Task { } } + private get requestOutcomeContext(): TaskRequestOutcomeContext { + return { + taskState: this.taskState, + messageStateHandler: this.messageStateHandler, + taskMessenger: this.taskMessenger, + api: this.api, + taskId: this.taskId, + checkpointManager: this.checkpointManager, + postStateToWebview: () => this.postStateToWebview(), + abortTask: () => this.abortTask(), + handleContextWindowExceededError: () => this.handleContextWindowExceededError(), + reinitExistingTaskFromId: () => this.reinitExistingTaskFromId(this.taskId), + attemptApiRequest: (previousApiReqIndex, lastApiReqIndex, shouldCompact) => + this.attemptApiRequest(previousApiReqIndex, lastApiReqIndex, shouldCompact), + recursivelyMakeDiracRequests: (userContent, includeFileDetails) => + this.recursivelyMakeDiracRequests(userContent, includeFileDetails), + handleEmptyAssistantResponse: (params) => this.handleEmptyAssistantResponse(params), + } + } + public canAcceptSteeringMessage(): boolean { return canAcceptSteeringMessage(this.steeringContext) } @@ -1182,187 +1202,7 @@ export class Task { providerId: string metricsManager: StreamingMetricsManager }): Promise { - const { error, model, providerId } = params - const diracError = DiracError.transform(error, model.id, providerId) - - if (diracError.isErrorType(DiracErrorType.ContextWindowExceeded)) { - await this.handleContextWindowExceededError() - const truncatedConversationHistory = this.messageStateHandler.getDiracMessages() - if (truncatedConversationHistory.length > 3) { - diracError.message = "Context window exceeded. Click retry to truncate the conversation and try again." - } - } - - const streamingFailedMessage = diracError.serialize() - - const lastApiReqStartedIndex = findLastIndex( - this.messageStateHandler.getDiracMessages(), - (m) => m.content.type === DiracMessageType.API_STATUS, - ) - if (lastApiReqStartedIndex !== -1) { - const diracMessages = this.messageStateHandler.getDiracMessages() - const msg = diracMessages[lastApiReqStartedIndex] - if (msg.content.type === DiracMessageType.API_STATUS) { - const currentApiReqInfo = { ...msg.content.status } - delete currentApiReqInfo.retryStatus - - await this.messageStateHandler.updateDiracMessage(lastApiReqStartedIndex, { - content: { - type: DiracMessageType.API_STATUS, - status: { - ...currentApiReqInfo, - streamingFailedMessage, - }, - }, - }) - } - } - - const isAuthError = diracError.isErrorType(DiracErrorType.Auth) - const isPaymentError = diracError.isErrorType(DiracErrorType.Payment) - - let response: DiracAskResponse - if (!isAuthError && !isPaymentError && this.taskState.apiErrorRetryAttempts < 3) { - this.taskState.apiErrorRetryAttempts++ - const delay = 2000 * 2 ** (this.taskState.apiErrorRetryAttempts - 1) - - await updateApiReqMsg({ - messageStateHandler: this.messageStateHandler, - lastApiReqIndex: lastApiReqStartedIndex, - inputTokens: 0, - reasoningTokens: 0, - outputTokens: 0, - cacheWriteTokens: 0, - cacheReadTokens: 0, - totalCost: undefined, - api: this.api, - cancelReason: "streaming_failed", - streamingFailedMessage, - }) - await this.messageStateHandler.saveDiracMessagesAndUpdateHistory() - await this.postStateToWebview() - - response = DiracAskResponse.APPROVE - const autoRetryCard = await this.taskMessenger.createCard({ - status: CardStatus.PENDING, - header: "API Error (Retrying)", - body: `API Error (attempt ${this.taskState.apiErrorRetryAttempts}/3). Retrying in ${delay / 1000}s...`, - }) - - const autoRetryApiReqIndex = findLastIndex( - this.messageStateHandler.getDiracMessages(), - (m) => m.content.type === DiracMessageType.API_STATUS, - ) - if (autoRetryApiReqIndex !== -1) { - const diracMessages = this.messageStateHandler.getDiracMessages() - const msg = diracMessages[autoRetryApiReqIndex] - if (msg.content.type === DiracMessageType.API_STATUS) { - const currentApiReqInfo = { ...msg.content.status } - delete currentApiReqInfo.streamingFailedMessage - await this.messageStateHandler.updateDiracMessage(autoRetryApiReqIndex, { - content: { - type: DiracMessageType.API_STATUS, - status: currentApiReqInfo, - }, - }) - } - } - - const deadline = Date.now() + delay - while (Date.now() < deadline && !this.taskState.abort) { - await setTimeoutPromise(Math.min(200, deadline - Date.now())) - } - // If the user aborted during the retry delay, stop retrying - if (this.taskState.abort) { - await autoRetryCard.update({ - header: "API Error (Cancelled)", - body: `API Error (attempt ${this.taskState.apiErrorRetryAttempts}/3). Cancelled.`, - }) - await autoRetryCard.finalize(CardStatus.CANCELLED) - throw new Error("Task instance aborted") - } - await autoRetryCard.update({ body: `API Error (attempt ${this.taskState.apiErrorRetryAttempts}/3). Retrying...` }) - await autoRetryCard.finalize(CardStatus.ERROR) - } else { - if (!isAuthError && !isPaymentError) { - await this.taskMessenger.createCard({ - status: CardStatus.ERROR, - header: "API Error (Retries Exhausted)", - body: `The API request failed after 3 attempts. ${diracError.toDisplayMessage()}`, - }) - } - if (isPaymentError) { - await this.taskMessenger.createCard({ - status: CardStatus.ERROR, - header: "API Error (Payment Required)", - body: diracError.toDisplayMessage(), - }) - } - if (isAuthError) { - await this.taskMessenger.createCard({ - status: CardStatus.ERROR, - header: "API Error (Authentication)", - body: diracError.toDisplayMessage(), - }) - } - this.taskState.status = TaskStatus.AWAITING_USER_INPUT - - const cardHandle = await this.taskMessenger.createCard({ - requireApproval: true, - header: "API Request Failed", - body: diracError.toDisplayMessage(), - actions: [ - { label: "Retry", value: DiracAskResponse.APPROVE, primary: true }, - { label: "Cancel", value: DiracAskResponse.REJECT }, - ], - }) - try { - const askResult = await cardHandle.waitForInteraction() - response = askResult.response - } catch (error) { - if (error instanceof ToolSkippedByUserMessage) { - await cardHandle.finalize(CardStatus.SKIPPED) - this.taskState.pendingUserMessage = error.userMessage - this.taskState.pendingUserImages = error.userImages - this.taskState.pendingUserFiles = error.userFiles - response = DiracAskResponse.APPROVE - } else { - throw error - } - } - if (response === DiracAskResponse.APPROVE) { - this.taskState.apiErrorRetryAttempts = 0 - } - } - - if (response !== DiracAskResponse.APPROVE) { - await this.abortTask() - await this.reinitExistingTaskFromId(this.taskId) - return false - } - - const manualRetryApiReqIndex = findLastIndex( - this.messageStateHandler.getDiracMessages(), - (m) => m.content.type === DiracMessageType.API_STATUS, - ) - if (manualRetryApiReqIndex !== -1) { - const diracMessages = this.messageStateHandler.getDiracMessages() - const msg = diracMessages[manualRetryApiReqIndex] - if (msg.content.type === DiracMessageType.API_STATUS) { - const currentApiReqInfo = { ...msg.content.status } - delete currentApiReqInfo.streamingFailedMessage - await this.messageStateHandler.updateDiracMessage(manualRetryApiReqIndex, { - content: { - type: DiracMessageType.API_STATUS, - status: currentApiReqInfo, - }, - }) - } - } - - await this.taskMessenger.upsertText("Retrying API request...") - - return true + return handleApiRequestError(this.requestOutcomeContext, params) } private async resetStreamingState(): Promise { @@ -1388,84 +1228,7 @@ export class Task { providerId: string model: { id: string } }): Promise { - if (params.assistantHasContent) { - this.taskState.askResponse = undefined - this.taskState.askResponseText = undefined - this.taskState.askResponseImages = undefined - this.taskState.askResponseFiles = undefined - this.taskState.status = TaskStatus.AWAITING_USER_INPUT - - await pWaitFor(() => this.taskState.userMessageContentReady) - const hasMutatingTools = this.taskState.assistantMessageContent.some( - (block) => block.type === "tool_use" && isMutatingTool(block.name), - ) - if (hasMutatingTools) { - await this.checkpointManager?.saveCheckpoint() - } - - const didToolUse = this.taskState.assistantMessageContent.some((block) => block.type === "tool_use") - if (this.taskState.didAttemptCompletion) { - this.taskState.completionCommitted = false - this.taskState.status = TaskStatus.COMPLETED - await this.postStateToWebview() - return true - } - const hitTokenLimit = - params.stopReason === "MAX_TOKENS" || params.stopReason === "max_tokens" || params.stopReason === "length" - - if (!didToolUse) { - this.taskState.userMessageContent.push({ - type: "text", - text: hitTokenLimit - ? "You reached the output token limit. Continue from where you stopped; restart an interrupted tool call, or call respond with operation 'complete' if finished." - : formatResponse.noToolsUsed(this.taskState.useNativeToolCalls), - }) - this.taskState.consecutiveMistakeCount++ - } - - this.taskState.apiErrorRetryAttempts = 0 - this.taskState.emptyResponseRetryAttempts = 0 - - if ( - this.taskState.pendingUserMessage !== undefined || - (this.taskState.pendingUserImages?.length ?? 0) > 0 || - (this.taskState.pendingUserFiles?.length ?? 0) > 0 - ) { - if (this.taskState.pendingUserMessage) { - this.taskState.userMessageContent.push({ - type: "text", - text: `\n${this.taskState.pendingUserMessage}\n`, - isUserInput: true, - }) - } - if (this.taskState.pendingUserImages?.length) { - this.taskState.userMessageContent.push(...formatResponse.imageBlocks(this.taskState.pendingUserImages)) - } - if (this.taskState.pendingUserFiles?.length) { - const fileContent = await processFilesIntoText(this.taskState.pendingUserFiles) - if (fileContent) { - this.taskState.userMessageContent.push({ type: "text", text: fileContent }) - } - } - this.taskState.pendingUserMessage = undefined - this.taskState.pendingUserImages = undefined - this.taskState.pendingUserFiles = undefined - } - - return await this.recursivelyMakeDiracRequests(this.taskState.userMessageContent) - } - const taskMetrics = params.metricsManager.getMetrics() - const shouldRetry = await this.handleEmptyAssistantResponse({ - modelInfo: params.modelInfo, - taskMetrics, - providerId: params.providerId, - model: params.model, - }) - if (shouldRetry === false) { - this.taskState.consecutiveMistakeCount = 0 - return await this.recursivelyMakeDiracRequests(params.userContent) - } - return true + return processStreamResult(this.requestOutcomeContext, params) } async *attemptApiRequest(previousApiReqIndex: number, lastApiReqIndex: number, shouldCompact?: boolean): ApiStream { From 73865c200f010e6dd3e3a2b6e5137bf0bdbf899b Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 9 Aug 2026 15:14:25 +0000 Subject: [PATCH 5/8] FB-15a: extract persistApiStopReason and handleMistakeLimitReached into helpers Co-Authored-By: Alexandros Salapatas --- src/core/task/TaskMistakeLimit.ts | 114 +++++++++++++++++++++++++ src/core/task/TaskRequestOutcome.ts | 24 +++++- src/core/task/index.ts | 125 +++------------------------- 3 files changed, 148 insertions(+), 115 deletions(-) create mode 100644 src/core/task/TaskMistakeLimit.ts diff --git a/src/core/task/TaskMistakeLimit.ts b/src/core/task/TaskMistakeLimit.ts new file mode 100644 index 00000000..e3ae1906 --- /dev/null +++ b/src/core/task/TaskMistakeLimit.ts @@ -0,0 +1,114 @@ +import { formatResponse } from "@core/formatResponse" +import { processFilesIntoText } from "@integrations/misc/extract-text" +import { showSystemNotification } from "@integrations/notifications" +import { CardStatus } from "@shared/ExtensionMessage" +import { DiracAskResponse } from "@shared/WebviewMessage" +import { DiracContent, type DiracUserContent } from "@shared/messages/content" +import type { StateManager } from "../storage/StateManager" +import type { TaskMessenger } from "./TaskMessenger" +import type { TaskState } from "./TaskState" +import { ToolSkippedByUserMessage } from "./tools/types/ToolSkippedByUserMessage" + +export interface TaskMistakeLimitContext { + taskState: TaskState + stateManager: StateManager + taskMessenger: TaskMessenger +} + +export async function handleMistakeLimitReached( + ctx: TaskMistakeLimitContext, + userContent: DiracContent[], +): Promise<{ didEndLoop: boolean; userContent: DiracContent[] }> { + if (ctx.taskState.consecutiveMistakeCount < ctx.stateManager.getGlobalSettingsKey("maxConsecutiveMistakes")) { + return { didEndLoop: false, userContent } + } + + // In yolo mode, don't wait for user input - fail the task + if (ctx.stateManager.getGlobalSettingsKey("yoloModeToggled")) { + const errorMessage = + `[YOLO MODE] Task failed: Too many consecutive mistakes (${ctx.taskState.consecutiveMistakeCount}). ` + + `The model may not be capable enough for this task. Consider using a more capable model.` + const card = await ctx.taskMessenger.createCard({ + status: CardStatus.ERROR, + header: "Task Failed", + body: errorMessage, + }) + await card.finalize(CardStatus.ERROR) + // End the task loop with failure + return { didEndLoop: true, userContent } // didEndLoop = true, signals task completion/failure + } + + const autoApprovalSettings = ctx.stateManager.getGlobalSettingsKey("autoApprovalSettings") + if (autoApprovalSettings.enableNotifications) { + showSystemNotification({ + subtitle: "Error", + message: "Dirac is having trouble. Would you like to continue the task?", + }) + } + + const cardHandle = await ctx.taskMessenger.createCard({ + header: "Mistake Limit Reached", + body: `Tool use failure. Can potentially be mitigated with some user guidance (e.g. "Try breaking down the task into smaller steps").`, + requireFeedback: true, + feedbackPlaceholder: "Provide guidance to Dirac...", + }) + let response: DiracAskResponse + let text: string | undefined + let images: string[] | undefined + let files: string[] | undefined + try { + const result = await cardHandle.waitForInteraction() + response = result.response + text = result.text + images = result.images + files = result.files + } catch (error) { + if (error instanceof ToolSkippedByUserMessage) { + await cardHandle.finalize(CardStatus.SKIPPED) + ctx.taskState.pendingUserMessage = error.userMessage + ctx.taskState.pendingUserImages = error.userImages + ctx.taskState.pendingUserFiles = error.userFiles + ctx.taskState.consecutiveMistakeCount = 0 + return { didEndLoop: false, userContent } + } + throw error + } + + await cardHandle.finalize(CardStatus.SUCCESS) + + if (response === DiracAskResponse.MESSAGE) { + // Display the user's message in the chat UI + await ctx.taskMessenger.upsertText(text || "", false, images, files, "user") + + // This userContent is for the *next* API call. + const feedbackUserContent: DiracUserContent[] = [] + feedbackUserContent.push({ + type: "text", + isUserInput: true, + text: formatResponse.tooManyMistakes(text), + }) + + if (images && images.length > 0) { + feedbackUserContent.push(...formatResponse.imageBlocks(images)) + } + + let fileContentString = "" + if (files && files.length > 0) { + fileContentString = await processFilesIntoText(files) + } + + if (fileContentString) { + feedbackUserContent.push({ + type: "text", + text: fileContentString, + }) + } + + userContent = feedbackUserContent + } + + ctx.taskState.consecutiveMistakeCount = 0 + ctx.taskState.apiErrorRetryAttempts = 0 + ctx.taskState.emptyResponseRetryAttempts = 0 + return { didEndLoop: false, userContent } +} diff --git a/src/core/task/TaskRequestOutcome.ts b/src/core/task/TaskRequestOutcome.ts index 9aa54d81..9bbd323e 100644 --- a/src/core/task/TaskRequestOutcome.ts +++ b/src/core/task/TaskRequestOutcome.ts @@ -7,10 +7,10 @@ import { processFilesIntoText } from "@integrations/misc/extract-text" import { DiracError, DiracErrorType } from "@services/error" import { findLastIndex } from "@shared/array" import { CardStatus, DiracMessageType, TaskStatus } from "@shared/ExtensionMessage" -import { DiracAskResponse } from "@shared/WebviewMessage" import { DiracContent, DiracTextContentBlock } from "@shared/messages/content" import type { DiracMessageModelInfo } from "@shared/messages/metrics" import { isMutatingTool } from "@shared/tools" +import { DiracAskResponse } from "@shared/WebviewMessage" import pWaitFor from "p-wait-for" import type { MessageStateHandler } from "./message-state" import type { StreamingMetricsManager } from "./StreamingMetricsManager" @@ -301,7 +301,7 @@ export async function processStreamResult( type: "text", text: `\n${ctx.taskState.pendingUserMessage}\n`, isUserInput: true, - } as DiracTextContentBlock) + } as DiracTextContentBlock) } if (ctx.taskState.pendingUserImages?.length) { ctx.taskState.userMessageContent.push(...formatResponse.imageBlocks(ctx.taskState.pendingUserImages)) @@ -332,3 +332,23 @@ export async function processStreamResult( } return true } + +export async function persistApiStopReason(ctx: TaskRequestOutcomeContext, stopReason?: string): Promise { + if (!stopReason) return + + const lastApiRequestIndex = findLastIndex( + ctx.messageStateHandler.getDiracMessages(), + (message) => message.content.type === DiracMessageType.API_STATUS, + ) + if (lastApiRequestIndex === -1) return + + const message = ctx.messageStateHandler.getDiracMessages()[lastApiRequestIndex] + if (message.content.type !== DiracMessageType.API_STATUS) return + + await ctx.messageStateHandler.updateDiracMessage(lastApiRequestIndex, { + content: { + type: DiracMessageType.API_STATUS, + status: { ...message.content.status, stopReason }, + }, + }) +} diff --git a/src/core/task/index.ts b/src/core/task/index.ts index 91d0ec66..4b4e4c88 100644 --- a/src/core/task/index.ts +++ b/src/core/task/index.ts @@ -20,7 +20,6 @@ import { ICheckpointManager } from "@integrations/checkpoints/types" import { DiffViewProvider } from "@integrations/editor/DiffViewProvider" import { FileEditProvider } from "@integrations/editor/FileEditProvider" import { processFilesIntoText } from "@integrations/misc/extract-text" -import { showSystemNotification } from "@integrations/notifications" import { type CommandExecutionOptions, type CommandExecutionResult, @@ -46,7 +45,6 @@ import { DiracStorageMessage, DiracTextContentBlock, DiracToolResponseContent, - DiracUserContent, removeProviderBoundaryMetadataFromMessage, } from "@shared/messages/content" import { DiracMessageModelInfo } from "@shared/messages/metrics" @@ -80,9 +78,15 @@ import { StreamingMetricsManager } from "./StreamingMetricsManager" import { StreamResponseHandler } from "./StreamResponseHandler" import { type SteeringClaim } from "./steering" import { TaskMessenger } from "./TaskMessenger" +import { handleMistakeLimitReached } from "./TaskMistakeLimit" import { type TaskPromptArtifactsContext, writePromptMetadataArtifacts } from "./TaskPromptArtifacts" import { buildApiRequestParams, type TaskRequestBuilderContext } from "./TaskRequestBuilder" -import { handleApiRequestError, processStreamResult, type TaskRequestOutcomeContext } from "./TaskRequestOutcome" +import { + handleApiRequestError, + persistApiStopReason, + processStreamResult, + type TaskRequestOutcomeContext, +} from "./TaskRequestOutcome" import { TaskState } from "./TaskState" import { appendQueuedSteeringToNextApiRequest, @@ -100,7 +104,6 @@ import { import { ToolExecutor } from "./ToolExecutor" import { DiracContext } from "./tools/context/DiracContext" import type { ToolSnapshotDirtyReason } from "./tools/runtime/ToolSnapshot" -import { ToolSkippedByUserMessage } from "./tools/types/ToolSkippedByUserMessage" import { extractProviderDomainFromUrl } from "./utils" export type ToolResponse = DiracToolResponseContent @@ -819,120 +822,16 @@ export class Task { } private async persistApiStopReason(stopReason?: string): Promise { - if (!stopReason) return - - const lastApiRequestIndex = findLastIndex( - this.messageStateHandler.getDiracMessages(), - (message) => message.content.type === DiracMessageType.API_STATUS, - ) - if (lastApiRequestIndex === -1) return - - const message = this.messageStateHandler.getDiracMessages()[lastApiRequestIndex] - if (message.content.type !== DiracMessageType.API_STATUS) return - - await this.messageStateHandler.updateDiracMessage(lastApiRequestIndex, { - content: { - type: DiracMessageType.API_STATUS, - status: { ...message.content.status, stopReason }, - }, - }) + return persistApiStopReason(this.requestOutcomeContext, stopReason) } private async handleMistakeLimitReached( userContent: DiracContent[], ): Promise<{ didEndLoop: boolean; userContent: DiracContent[] }> { - if (this.taskState.consecutiveMistakeCount < this.stateManager.getGlobalSettingsKey("maxConsecutiveMistakes")) { - return { didEndLoop: false, userContent } - } - - // In yolo mode, don't wait for user input - fail the task - if (this.stateManager.getGlobalSettingsKey("yoloModeToggled")) { - const errorMessage = - `[YOLO MODE] Task failed: Too many consecutive mistakes (${this.taskState.consecutiveMistakeCount}). ` + - `The model may not be capable enough for this task. Consider using a more capable model.` - const card = await this.taskMessenger.createCard({ - status: CardStatus.ERROR, - header: "Task Failed", - body: errorMessage, - }) - await card.finalize(CardStatus.ERROR) - // End the task loop with failure - return { didEndLoop: true, userContent } // didEndLoop = true, signals task completion/failure - } - - const autoApprovalSettings = this.stateManager.getGlobalSettingsKey("autoApprovalSettings") - if (autoApprovalSettings.enableNotifications) { - showSystemNotification({ - subtitle: "Error", - message: "Dirac is having trouble. Would you like to continue the task?", - }) - } - - const cardHandle = await this.taskMessenger.createCard({ - header: "Mistake Limit Reached", - body: `Tool use failure. Can potentially be mitigated with some user guidance (e.g. "Try breaking down the task into smaller steps").`, - requireFeedback: true, - feedbackPlaceholder: "Provide guidance to Dirac...", - }) - let response: DiracAskResponse - let text: string | undefined - let images: string[] | undefined - let files: string[] | undefined - try { - const result = await cardHandle.waitForInteraction() - response = result.response - text = result.text - images = result.images - files = result.files - } catch (error) { - if (error instanceof ToolSkippedByUserMessage) { - await cardHandle.finalize(CardStatus.SKIPPED) - this.taskState.pendingUserMessage = error.userMessage - this.taskState.pendingUserImages = error.userImages - this.taskState.pendingUserFiles = error.userFiles - this.taskState.consecutiveMistakeCount = 0 - return { didEndLoop: false, userContent } - } - throw error - } - - await cardHandle.finalize(CardStatus.SUCCESS) - - if (response === DiracAskResponse.MESSAGE) { - // Display the user's message in the chat UI - await this.taskMessenger.upsertText(text || "", false, images, files, "user") - - // This userContent is for the *next* API call. - const feedbackUserContent: DiracUserContent[] = [] - feedbackUserContent.push({ - type: "text", - isUserInput: true, - text: formatResponse.tooManyMistakes(text), - }) - - if (images && images.length > 0) { - feedbackUserContent.push(...formatResponse.imageBlocks(images)) - } - - let fileContentString = "" - if (files && files.length > 0) { - fileContentString = await processFilesIntoText(files) - } - - if (fileContentString) { - feedbackUserContent.push({ - type: "text", - text: fileContentString, - }) - } - - userContent = feedbackUserContent - } - - this.taskState.consecutiveMistakeCount = 0 - this.taskState.apiErrorRetryAttempts = 0 - this.taskState.emptyResponseRetryAttempts = 0 - return { didEndLoop: false, userContent } + return handleMistakeLimitReached( + { taskState: this.taskState, stateManager: this.stateManager, taskMessenger: this.taskMessenger }, + userContent, + ) } async loadContext( From ff7f3c419472c54c54f488114e530d4c7fbe7dfc Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 9 Aug 2026 15:17:58 +0000 Subject: [PATCH 6/8] FB-15a: extract waitForFollowUp and submitCardResponse into TaskUserInput.ts Co-Authored-By: Alexandros Salapatas --- src/core/task/TaskUserInput.ts | 77 ++++++++++++++++++++++++++++++++++ src/core/task/index.ts | 57 ++----------------------- 2 files changed, 81 insertions(+), 53 deletions(-) create mode 100644 src/core/task/TaskUserInput.ts diff --git a/src/core/task/TaskUserInput.ts b/src/core/task/TaskUserInput.ts new file mode 100644 index 00000000..c8256b0c --- /dev/null +++ b/src/core/task/TaskUserInput.ts @@ -0,0 +1,77 @@ +import { formatResponse } from "@core/formatResponse" +import { processFilesIntoText } from "@integrations/misc/extract-text" +import { Logger } from "@shared/services/Logger" +import { TaskStatus } from "@shared/ExtensionMessage" +import { DiracAskResponse } from "@shared/WebviewMessage" +import { DiracContent } from "@shared/messages/content" +import pWaitFor from "p-wait-for" +import type { TaskState } from "./TaskState" + +export interface TaskUserInputContext { + taskState: TaskState +} + +export async function waitForFollowUp(ctx: TaskUserInputContext): Promise { + ctx.taskState.status = TaskStatus.AWAITING_USER_INPUT + + const messageTs = Date.now() + ctx.taskState.lastMessageTs = messageTs + + await pWaitFor( + () => { + return ctx.taskState.askResponse !== undefined || ctx.taskState.lastMessageTs !== messageTs || ctx.taskState.abort + }, + { interval: 100 }, + ) + + if (ctx.taskState.abort || ctx.taskState.lastMessageTs !== messageTs) { + return undefined + } + + const text = ctx.taskState.askResponseText || "" + const images = ctx.taskState.askResponseImages as string[] | undefined + const files = ctx.taskState.askResponseFiles as string[] | undefined + + const userContent: DiracContent[] = [{ type: "text", text: `\n${text}\n`, isUserInput: true }] + if (images && images.length > 0) { + userContent.push(...formatResponse.imageBlocks(images)) + } + if (files && files.length > 0) { + const fileContentString = await processFilesIntoText(files) + if (fileContentString) { + userContent.push({ type: "text", text: fileContentString }) + } + } + + return userContent +} + +export async function submitCardResponse( + ctx: TaskUserInputContext, + params: { + cardId: string + response: DiracAskResponse | string + text?: string + images?: string[] + files?: string[] + value?: string + }, +): Promise { + const { cardId, response, text, images, files, value } = params + if (cardId && ctx.taskState.lastWaitingCardId !== cardId) { + Logger.warn(`[Task] Received response for card ${cardId}, but waiting for ${ctx.taskState.lastWaitingCardId}`) + return + } + const isStandardResponse = Object.values(DiracAskResponse).includes(response as DiracAskResponse) + ctx.taskState.askResponse = isStandardResponse ? (response as DiracAskResponse) : undefined + ctx.taskState.askResponseText = text + ctx.taskState.askResponseImages = images + ctx.taskState.askResponseFiles = files + ctx.taskState.askResponseAction = response as string + ctx.taskState.askResponseValue = value + // When user sends a text message while a card is awaiting approval, + // signal that the tool should be skipped and forward to LLM. + if (response === DiracAskResponse.MESSAGE && text && ctx.taskState.status !== TaskStatus.CANCELLED) { + ctx.taskState.didRejectTool = true + } +} diff --git a/src/core/task/index.ts b/src/core/task/index.ts index 4b4e4c88..b40fae94 100644 --- a/src/core/task/index.ts +++ b/src/core/task/index.ts @@ -4,7 +4,7 @@ import { ContextManager } from "@core/context/context-management/ContextManager" import { EnvironmentContextTracker } from "@core/context/context-tracking/EnvironmentContextTracker" import { FileContextTracker } from "@core/context/context-tracking/FileContextTracker" import { ModelContextTracker } from "@core/context/context-tracking/ModelContextTracker" -import { formatResponse } from "@core/formatResponse" + import { DiracIgnoreController } from "@core/ignore/DiracIgnoreController" import { recordSuccessfulModelProviderPreset } from "@core/models/modelProviderPresets" import { CommandPermissionController } from "@core/permissions" @@ -19,7 +19,6 @@ import { buildCheckpointManager, shouldUseMultiRoot } from "@integrations/checkp import { ICheckpointManager } from "@integrations/checkpoints/types" import { DiffViewProvider } from "@integrations/editor/DiffViewProvider" import { FileEditProvider } from "@integrations/editor/FileEditProvider" -import { processFilesIntoText } from "@integrations/misc/extract-text" import { type CommandExecutionOptions, type CommandExecutionResult, @@ -56,7 +55,6 @@ import { type Mode } from "@shared/storage/types" import { DiracAskResponse } from "@shared/WebviewMessage" import { isLocalModel, isParallelToolCallingEnabled } from "@utils/model-utils" import Mutex from "p-mutex" -import pWaitFor from "p-wait-for" import { ulid } from "ulid" import { getErrorMessage } from "@/shared/errors" import { type SkillMetadata } from "@/shared/skills" @@ -105,6 +103,7 @@ import { ToolExecutor } from "./ToolExecutor" import { DiracContext } from "./tools/context/DiracContext" import type { ToolSnapshotDirtyReason } from "./tools/runtime/ToolSnapshot" import { extractProviderDomainFromUrl } from "./utils" +import { submitCardResponse, waitForFollowUp } from "./TaskUserInput" export type ToolResponse = DiracToolResponseContent @@ -855,40 +854,7 @@ export class Task { } private async waitForFollowUp(): Promise { - this.taskState.status = TaskStatus.AWAITING_USER_INPUT - - const messageTs = Date.now() - this.taskState.lastMessageTs = messageTs - - await pWaitFor( - () => { - return ( - this.taskState.askResponse !== undefined || this.taskState.lastMessageTs !== messageTs || this.taskState.abort - ) - }, - { interval: 100 }, - ) - - if (this.taskState.abort || this.taskState.lastMessageTs !== messageTs) { - return undefined - } - - const text = this.taskState.askResponseText || "" - const images = this.taskState.askResponseImages as string[] | undefined - const files = this.taskState.askResponseFiles as string[] | undefined - - const userContent: DiracContent[] = [{ type: "text", text: `\n${text}\n`, isUserInput: true }] - if (images && images.length > 0) { - userContent.push(...formatResponse.imageBlocks(images)) - } - if (files && files.length > 0) { - const fileContentString = await processFilesIntoText(files) - if (fileContentString) { - userContent.push({ type: "text", text: fileContentString }) - } - } - - return userContent + return waitForFollowUp({ taskState: this.taskState }) } /** Persist a project-scoped tool permission rule for an ACP “always” decision. */ @@ -915,22 +881,7 @@ export class Task { value?: string, ) { await this.withStateLock(async () => { - if (cardId && this.taskState.lastWaitingCardId !== cardId) { - Logger.warn(`[Task] Received response for card ${cardId}, but waiting for ${this.taskState.lastWaitingCardId}`) - return - } - const isStandardResponse = Object.values(DiracAskResponse).includes(response as DiracAskResponse) - this.taskState.askResponse = isStandardResponse ? (response as DiracAskResponse) : undefined - this.taskState.askResponseText = text - this.taskState.askResponseImages = images - this.taskState.askResponseFiles = files - this.taskState.askResponseAction = response as string - this.taskState.askResponseValue = value - // When user sends a text message while a card is awaiting approval, - // signal that the tool should be skipped and forward to LLM. - if (response === DiracAskResponse.MESSAGE && text && this.taskState.status !== TaskStatus.CANCELLED) { - this.taskState.didRejectTool = true - } + return submitCardResponse({ taskState: this.taskState }, { cardId, response, text, images, files, value }) }) } From 00bd674d740f887dacce7188c3e757712bc29acc Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 9 Aug 2026 15:26:44 +0000 Subject: [PATCH 7/8] FB-15a: extract attemptApiRequest and recursivelyMakeDiracRequests into TaskRequestLoop.ts Co-Authored-By: Alexandros Salapatas --- src/core/task/TaskRequestLoop.ts | 549 +++++++++++++++++++++++++++++ src/core/task/index.ts | 582 ++----------------------------- 2 files changed, 581 insertions(+), 550 deletions(-) create mode 100644 src/core/task/TaskRequestLoop.ts diff --git a/src/core/task/TaskRequestLoop.ts b/src/core/task/TaskRequestLoop.ts new file mode 100644 index 00000000..61d3775a --- /dev/null +++ b/src/core/task/TaskRequestLoop.ts @@ -0,0 +1,549 @@ +import { ApiStream } from "@core/api/transform/stream" +import { recordSuccessfulModelProviderPreset } from "@core/models/modelProviderPresets" +import { ErrorService } from "@services/error" +import { telemetryService } from "@services/telemetry" +import type { ApiProvider } from "@shared/api" +import { findLastIndex } from "@shared/array" +import type { DiracApiReqCancelReason, DiracMessageContent } from "@shared/ExtensionMessage" +import { CardStatus, DiracMessageType, TaskStatus } from "@shared/ExtensionMessage" +import { Session } from "@shared/services/Session" +import { isLocalModel } from "@utils/model-utils" +import type { DiffViewProvider } from "@integrations/editor/DiffViewProvider" +import type { LocalConversationCompaction } from "./LocalConversationCompaction" + +import type { ModelContextTracker } from "@core/context/context-tracking/ModelContextTracker" +import type { ResponseProcessor } from "./ResponseProcessor" +import { StreamChunkCoordinator } from "./StreamChunkCoordinator" +import { StreamingMetricsManager } from "./StreamingMetricsManager" +import type { StreamResponseHandler } from "./StreamResponseHandler" + +import { removeProviderBoundaryMetadataFromMessage } from "@shared/messages/content" +import type { DiracContent, DiracStorageMessage } from "@shared/messages/content" +import type { DiracMessageModelInfo } from "@shared/messages/metrics" +import { Logger } from "@shared/services/Logger" +import { buildApiRequestParams, type TaskRequestBuilderContext } from "./TaskRequestBuilder" +import { + handleApiRequestError, + persistApiStopReason, + processStreamResult, + type TaskRequestOutcomeContext, +} from "./TaskRequestOutcome" +import { + appendQueuedSteeringToNextApiRequest, + appendQueuedSteeringToUserContent, + rollbackSteeringClaim, + settleConsumedSteeringClaim, + type TaskSteeringContext, +} from "./TaskSteering" +import type { ApiConversationManager } from "./ApiConversationManager" + +export interface TaskRequestLoopContext extends TaskRequestBuilderContext, TaskRequestOutcomeContext { + steeringContext: TaskSteeringContext + handleMistakeLimitReached: (userContent: DiracContent[]) => Promise<{ didEndLoop: boolean; userContent: DiracContent[] }> + enqueuePreRequestSteeringMessages: () => Promise + resetStreamingState: () => Promise + initializeCheckpoints: (isFirstRequest: boolean) => Promise + determineContextCompaction: (previousApiReqIndex: number) => Promise + localConversationCompaction: LocalConversationCompaction + responseProcessor: ResponseProcessor + streamHandler: StreamResponseHandler + modelContextTracker: ModelContextTracker + diffViewProvider: DiffViewProvider + ulid: string +} + +export async function* attemptApiRequest( + ctx: TaskRequestLoopContext, + previousApiReqIndex: number, + lastApiReqIndex: number, + shouldCompact?: boolean, +): ApiStream { + const { systemPrompt, toolSnapshot, contextManagementMetadata, providerInfo } = await buildApiRequestParams(ctx, { + previousApiReqIndex, + shouldCompact, + }) + const { model, providerId } = providerInfo + + const metricsManager = new StreamingMetricsManager(ctx.messageStateHandler, lastApiReqIndex, ctx.api) + + const finalizeApiReqMsg = async (cancelReason?: DiracApiReqCancelReason, streamingFailedMessage?: string) => { + await metricsManager.updateApiReqMsgFromMetrics(cancelReason, streamingFailedMessage) + await ctx.messageStateHandler.updateDiracMessage(lastApiReqIndex, {}) + ctx.taskState.isApiRequestActive = false + ctx.taskState.activeVoiceStreamId = undefined + } + + const abortStream = async (cancelReason: DiracApiReqCancelReason, streamingFailedMessage?: string) => { + ctx.taskState.didFinishAbortingStream = true + await finalizeApiReqMsg(cancelReason, streamingFailedMessage) + ctx.taskState.isApiRequestActive = false + ctx.taskState.activeVoiceStreamId = undefined + } + + await appendQueuedSteeringToNextApiRequest(ctx.steeringContext, contextManagementMetadata.truncatedConversationHistory) + + const providerDispatch = await ctx.apiConversationManager.prepareProviderConversationDispatch({ + systemPrompt, + tools: toolSnapshot.nativeTools, + truncatedMessages: contextManagementMetadata.truncatedConversationHistory as DiracStorageMessage[], + providerId, + modelId: model.id, + }) + + if (ctx.taskState.abort) throw new Error("Task instance aborted") + + const stream = ctx.api.createMessage( + systemPrompt, + providerDispatch.messages.map(removeProviderBoundaryMetadataFromMessage), + toolSnapshot.nativeTools, + providerDispatch.options, + ) + const iterator = stream[Symbol.asyncIterator]() + + try { + ctx.taskState.status = TaskStatus.WAITING_FOR_API + + ctx.taskState.isWaitingForFirstChunk = true + const firstChunk = await iterator.next() + ctx.taskState.isWaitingForFirstChunk = false + + if (firstChunk.done) { + await finalizeApiReqMsg() + return + } + + yield firstChunk.value + + for await (const chunk of iterator) { + if (ctx.taskState.abort) { + await abortStream("user_cancelled") + return + } + + if (chunk.type === "usage") { + metricsManager.updateFromChunk(chunk) + yield chunk + continue + } + + yield chunk + } + + recordSuccessfulModelProviderPreset( + ctx.stateManager, + providerId as ApiProvider, + model.id, + model.info, + providerInfo.mode, + ) + await finalizeApiReqMsg() + } catch (error) { + const shouldRetry = await handleApiRequestError(ctx, { + error, + previousApiReqIndex, + lastApiReqIndex, + shouldCompact, + model, + providerId, + metricsManager, + }) + if (shouldRetry) { + yield* attemptApiRequest(ctx, previousApiReqIndex, lastApiReqIndex, shouldCompact) + } + return + } +} + +export async function recursivelyMakeDiracRequests( + ctx: TaskRequestLoopContext, + userContent: DiracContent[], + includeFileDetails = false, +): Promise { + ctx.taskState.status = TaskStatus.PREPARING + + if (ctx.taskState.abort) { + throw new Error("Task instance aborted") + } + await ctx.enqueuePreRequestSteeringMessages() + + const { model, providerId, customPrompt, mode } = ctx.getCurrentProviderInfo() + if (providerId && model.id) { + try { + await ctx.modelContextTracker.recordModelUsage(providerId, model.id, mode) + } catch (error) { + Logger.error("Failed to record model usage:", error) + } + } + + const modelInfo: DiracMessageModelInfo = { + modelId: model.id, + providerId: providerId, + mode: mode, + } + + const mistakeResult = await ctx.handleMistakeLimitReached(userContent) + if (mistakeResult.didEndLoop) { + return true + } + userContent = mistakeResult.userContent + + const previousApiReqIndex = findLastIndex( + ctx.messageStateHandler.getDiracMessages(), + (m) => m.content.type === DiracMessageType.API_STATUS + ) + const isFirstRequest = + ctx.messageStateHandler.getDiracMessages().filter((m) => m.content.type === DiracMessageType.API_STATUS).length === 0 + + await ctx.initializeCheckpoints(isFirstRequest) + + const useCompactPrompt = customPrompt === "compact" && isLocalModel(ctx.getCurrentProviderInfo()) + let shouldCompact = await ctx.determineContextCompaction(previousApiReqIndex) + if (shouldCompact && ctx.localConversationCompaction.isAvailable()) { + const continuation = await ctx.localConversationCompaction.run({ + source: "automatic", + triggerApiRequestIndex: previousApiReqIndex, + }) + if (!continuation) return true + + const compactedContext: DiracContent[] = [{ type: "text", text: continuation }] + if (ctx.taskState.pinnedContext) { + compactedContext.push({ type: "text", text: ctx.taskState.pinnedContext }) + } + userContent = [...compactedContext, ...userContent] + shouldCompact = false + } + const steeringClaim = await appendQueuedSteeringToUserContent(ctx.steeringContext, userContent) + + ctx.taskState.status = TaskStatus.BUILDING_REQUEST + + let apiRequestData: Awaited> + let steeringClaimConsumed = false + try { + apiRequestData = await ctx.apiConversationManager.prepareApiRequest({ + userContent, + shouldCompact, + includeFileDetails, + useCompactPrompt, + previousApiReqIndex, + isFirstRequest, + providerId, + modelId: model.id, + mode: modelInfo.mode, + afterUserContentPersisted: async () => { + steeringClaimConsumed = true + if (!steeringClaim) return + await settleConsumedSteeringClaim(ctx.steeringContext, steeringClaim) + }, + }) + if (steeringClaim && !steeringClaimConsumed) { + if (apiRequestData.didConsumeUserContent) { + steeringClaimConsumed = true + await settleConsumedSteeringClaim(ctx.steeringContext, steeringClaim) + } else { + await rollbackSteeringClaim(ctx.steeringContext, steeringClaim.id) + } + } + } catch (error) { + if (steeringClaim && !steeringClaimConsumed) await rollbackSteeringClaim(ctx.steeringContext, steeringClaim.id) + throw error + } + userContent = apiRequestData.userContent + const lastApiReqIndex = apiRequestData.lastApiReqIndex + + if (apiRequestData.isDirectResponse) { + if (apiRequestData.directResponseText) { + await ctx.taskMessenger.upsertText(apiRequestData.directResponseText) + } + return true + } + + try { + const metricsManager = new StreamingMetricsManager(ctx.messageStateHandler, lastApiReqIndex, ctx.api) + let didFinalizeApiReqMsg = false + let usageChunkSideEffectsQueue = Promise.resolve() + + const queueUsageChunkSideEffects = ( + usageInputTokens: number, + usageOutputTokens: number, + chunkOptions?: { cacheWriteTokens?: number; cacheReadTokens?: number; totalCost?: number; stopReason?: string }, + ) => { + usageChunkSideEffectsQueue = usageChunkSideEffectsQueue.then(async () => { + if (didFinalizeApiReqMsg || ctx.taskState.abort) { + return + } + + await metricsManager.updateApiReqMsgFromMetrics() + await ctx.postStateToWebview() + await telemetryService.captureTokenUsage( + ctx.ulid, + usageInputTokens, + usageOutputTokens, + providerId, + model.id, + chunkOptions, + ) + }) + } + + const finalizeApiReqMsg = async (cancelReason?: DiracApiReqCancelReason, streamingFailedMessage?: string) => { + didFinalizeApiReqMsg = true + await usageChunkSideEffectsQueue + await metricsManager.updateApiReqMsgFromMetrics(cancelReason, streamingFailedMessage) + + const metrics = metricsManager.getMetrics() + ctx.taskState.totalInputTokens += metrics.inputTokens + ctx.taskState.totalOutputTokens += metrics.outputTokens + ctx.taskState.totalReasoningTokens += metrics.reasoningTokens + ctx.taskState.totalCacheWriteTokens += metrics.cacheWriteTokens + ctx.taskState.totalCacheReadTokens += metrics.cacheReadTokens + ctx.taskState.totalCost += metricsManager.getTotalCost() + + const currentApiReqIndex = findLastIndex( + ctx.messageStateHandler.getDiracMessages(), + (m) => m.content.type === DiracMessageType.API_STATUS, + ) + if (currentApiReqIndex !== -1) { + ctx.taskState.isApiRequestActive = false + ctx.taskState.activeVoiceStreamId = undefined + } + } + + const abortStream = async (cancelReason: DiracApiReqCancelReason, streamingFailedMessage?: string) => { + Session.get().finalizeRequest() + + if (ctx.diffViewProvider.isEditing) { + await ctx.diffViewProvider.revertChanges() + } + + ctx.taskState.isApiRequestActive = false + ctx.taskState.activeVoiceStreamId = undefined + await finalizeApiReqMsg(cancelReason, streamingFailedMessage) + await ctx.messageStateHandler.saveDiracMessagesAndUpdateHistory() + + const metrics = metricsManager.getMetrics() + await ctx.messageStateHandler.addToApiConversationHistory({ + role: "assistant", + content: [ + { + type: "text", + text: + assistantMessage + + `\n\n[${ + cancelReason === "streaming_failed" + ? "Response interrupted by API Error" + : "Response interrupted by user" + }]`, + }, + ], + modelInfo, + metrics: { + tokens: { + prompt: metrics.inputTokens, + completion: metrics.outputTokens, + cached: (metrics.cacheWriteTokens ?? 0) + (metrics.cacheReadTokens ?? 0), + }, + cost: metrics.totalCost, + }, + ts: Date.now(), + }) + + telemetryService.captureConversationTurnEvent( + ctx.ulid, + providerId, + modelInfo.modelId, + "assistant", + modelInfo.mode, + undefined, + ctx.taskState.useNativeToolCalls, + ) + + ctx.taskState.didFinishAbortingStream = true + } + + await ctx.resetStreamingState() + + const { toolUseHandler, reasonsHandler } = ctx.streamHandler.getHandlers() + const stream = attemptApiRequest(ctx, previousApiReqIndex, lastApiReqIndex, shouldCompact) + + let assistantMessageId = "" + let assistantMessage = "" + let assistantTextOnly = "" + let assistantTextSignature: string | undefined + + let didReceiveUsageChunk = false + let stopReason: string | undefined + let didFinalizeReasoningForUi = false + + const finalizePendingReasoningMessage = async (thinking: string): Promise => { + const activeVoiceStreamId = ctx.taskState.activeVoiceStreamId + if (!activeVoiceStreamId) { + return false + } + + const messages = ctx.messageStateHandler.getDiracMessages() + const pendingReasoningIndex = messages.findIndex((m) => m.id === activeVoiceStreamId) + + if (pendingReasoningIndex !== -1) { + const msg = messages[pendingReasoningIndex] + if (msg.content.type === DiracMessageType.MARKDOWN && msg.content.isReasoning) { + await ctx.messageStateHandler.updateDiracMessage(pendingReasoningIndex, { + content: { type: DiracMessageType.MARKDOWN, content: thinking, isReasoning: true }, + }) + const completedReasoning = ctx.messageStateHandler.getDiracMessages()[pendingReasoningIndex] + if (completedReasoning) { + await ctx.postStateToWebview() + } + ctx.taskState.activeVoiceStreamId = undefined + return true + } + } + return false + } + + Session.get().startApiCall() + ctx.taskState.isApiRequestActive = true + let streamCoordinator: StreamChunkCoordinator | undefined + + try { + streamCoordinator = new StreamChunkCoordinator(stream, { + onUsageChunk: (chunk) => { + ctx.streamHandler.setRequestId(chunk.id) + didReceiveUsageChunk = true + metricsManager.updateFromChunk(chunk) + stopReason = chunk.stopReason ?? stopReason + queueUsageChunkSideEffects(chunk.inputTokens, chunk.outputTokens, { + cacheWriteTokens: chunk.cacheWriteTokens, + cacheReadTokens: chunk.cacheReadTokens, + totalCost: chunk.totalCost, + stopReason: chunk.stopReason, + }) + }, + }) + + const streamResult = await ctx.responseProcessor.consumeStream(streamCoordinator, { + abortStream, + finalizePendingReasoningMessage, + apiAbort: () => ctx.api.abort?.(), + }) + + assistantMessage = streamResult.assistantMessage + assistantTextOnly = streamResult.assistantTextOnly + assistantTextSignature = streamResult.assistantTextSignature + assistantMessageId = streamResult.assistantMessageId + didFinalizeReasoningForUi = streamResult.didFinalizeReasoningForUi + const shouldInterruptStream = streamResult.shouldInterruptStream + + if (shouldInterruptStream) { + await streamCoordinator.stop() + } else { + await streamCoordinator.waitForCompletion() + } + await usageChunkSideEffectsQueue + + if (!ctx.taskState.abort && !didFinalizeReasoningForUi) { + const finalReasoning = reasonsHandler.getCurrentReasoning() + if (finalReasoning?.thinking) { + await finalizePendingReasoningMessage(finalReasoning.thinking) + didFinalizeReasoningForUi = true + } + } + } catch (error) { + await streamCoordinator?.stop() + if (ctx.taskState.abort || ctx.taskState.abandoned) { + return true + } + + const diracError = ErrorService.get().toDiracError(error, ctx.api.getModel().id) + const errorMessage = diracError.serialize() + await ctx.abortTask() + await abortStream("streaming_failed", errorMessage) + await ctx.reinitExistingTaskFromId() + return true + } finally { + Session.get().endApiCall() + } + + if (!didReceiveUsageChunk) { + const apiStreamUsage = await ctx.api.getApiStreamUsage?.() + if (apiStreamUsage) { + metricsManager.updateFromChunk(apiStreamUsage) + queueUsageChunkSideEffects(apiStreamUsage.inputTokens, apiStreamUsage.outputTokens, { + cacheWriteTokens: apiStreamUsage.cacheWriteTokens, + cacheReadTokens: apiStreamUsage.cacheReadTokens, + totalCost: apiStreamUsage.totalCost, + stopReason: apiStreamUsage.stopReason, + }) + } + } + + const autoRetryApiReqIndex = findLastIndex( + ctx.messageStateHandler.getDiracMessages(), + (m) => m.content.type === DiracMessageType.API_STATUS, + ) + if (autoRetryApiReqIndex !== -1) { + const diracMessages = ctx.messageStateHandler.getDiracMessages() + const msg = diracMessages[autoRetryApiReqIndex] + if (msg.content.type === DiracMessageType.API_STATUS) { + const content = msg.content as Extract + const currentApiReqInfo = { ...content.status } + delete currentApiReqInfo.retryStatus + await ctx.messageStateHandler.updateDiracMessage(autoRetryApiReqIndex, { + content: { + type: DiracMessageType.API_STATUS, + status: currentApiReqInfo, + }, + }) + } + } + + await finalizeApiReqMsg() + await persistApiStopReason(ctx, stopReason) + await ctx.messageStateHandler.saveDiracMessagesAndUpdateHistory() + await ctx.postStateToWebview() + + if (ctx.taskState.abort) { + throw new Error("Dirac instance aborted") + } + + const assistantHasContent = await ctx.responseProcessor.routeAssistantResponse({ + assistantMessage, + assistantTextOnly, + assistantTextSignature, + assistantMessageId, + providerId, + modelId: model.id, + mode: modelInfo.mode, + taskMetrics: metricsManager.getMetrics(), + modelInfo, + toolUseHandler, + }) + + return await processStreamResult(ctx, { + assistantHasContent, + stopReason, + userContent, + metricsManager, + modelInfo, + providerId, + model, + }) + } catch (error) { + if (ctx.taskState.abort) { + // User-initiated abort — not a fatal error, no card needed + return true + } + const diracError = ErrorService.get().toDiracError(error) + Logger.error("[Task] Fatal error in task loop:", diracError.serialize()) + try { + const card = await ctx.taskMessenger.createCard({ + status: CardStatus.ERROR, + header: "Task Error", + body: `The task encountered an unexpected error and had to stop.\n\n${diracError.toDisplayMessage()}`, + }) + await card.finalize(CardStatus.ERROR) + } catch (sayError) { + Logger.error("[Task] Failed to emit error message:", sayError) + } + return true + } +} diff --git a/src/core/task/index.ts b/src/core/task/index.ts index b40fae94..c131c9ca 100644 --- a/src/core/task/index.ts +++ b/src/core/task/index.ts @@ -6,7 +6,7 @@ import { FileContextTracker } from "@core/context/context-tracking/FileContextTr import { ModelContextTracker } from "@core/context/context-tracking/ModelContextTracker" import { DiracIgnoreController } from "@core/ignore/DiracIgnoreController" -import { recordSuccessfulModelProviderPreset } from "@core/models/modelProviderPresets" + import { CommandPermissionController } from "@core/permissions" import type { SlashCommandDirectAction } from "@core/slash-commands" import { createDefaultTextCondensationTemplateRegistry } from "@core/text-condensation/templates" @@ -30,13 +30,11 @@ import { import { ITerminalManager } from "@integrations/terminal/types" import { BrowserSession } from "@services/browser/BrowserSession" import { UrlContentFetcher } from "@services/browser/UrlContentFetcher" -import { ErrorService } from "@services/error" - import { telemetryService } from "@services/telemetry" import { ApiConfiguration } from "@shared/api" -import { findLastIndex } from "@shared/array" + import { getExtensionSourceDir } from "@shared/dirac/constants" -import { CardStatus, DiracApiReqCancelReason, DiracMessageContent, DiracMessageType, TaskStatus } from "@shared/ExtensionMessage" +import { TaskStatus } from "@shared/ExtensionMessage" import { HistoryItem } from "@shared/HistoryItem" import { @@ -44,16 +42,15 @@ import { DiracStorageMessage, DiracTextContentBlock, DiracToolResponseContent, - removeProviderBoundaryMetadataFromMessage, } from "@shared/messages/content" -import { DiracMessageModelInfo } from "@shared/messages/metrics" + import { ShowMessageType } from "@shared/proto/index.host" import { Logger } from "@shared/services/Logger" -import { Session } from "@shared/services/Session" + import { type Mode } from "@shared/storage/types" import { DiracAskResponse } from "@shared/WebviewMessage" -import { isLocalModel, isParallelToolCallingEnabled } from "@utils/model-utils" +import { isParallelToolCallingEnabled } from "@utils/model-utils" import Mutex from "p-mutex" import { ulid } from "ulid" import { getErrorMessage } from "@/shared/errors" @@ -71,20 +68,20 @@ import { LifecycleManager } from "./LifecycleManager" import { LocalConversationCompaction } from "./LocalConversationCompaction" import { MessageStateHandler } from "./message-state" import { ResponseProcessor } from "./ResponseProcessor" -import { StreamChunkCoordinator } from "./StreamChunkCoordinator" + import { StreamingMetricsManager } from "./StreamingMetricsManager" import { StreamResponseHandler } from "./StreamResponseHandler" import { type SteeringClaim } from "./steering" import { TaskMessenger } from "./TaskMessenger" import { handleMistakeLimitReached } from "./TaskMistakeLimit" import { type TaskPromptArtifactsContext, writePromptMetadataArtifacts } from "./TaskPromptArtifacts" -import { buildApiRequestParams, type TaskRequestBuilderContext } from "./TaskRequestBuilder" +import { type TaskRequestBuilderContext } from "./TaskRequestBuilder" import { handleApiRequestError, persistApiStopReason, - processStreamResult, type TaskRequestOutcomeContext, } from "./TaskRequestOutcome" +import { attemptApiRequest, recursivelyMakeDiracRequests, type TaskRequestLoopContext } from "./TaskRequestLoop" import { TaskState } from "./TaskState" import { appendQueuedSteeringToNextApiRequest, @@ -210,7 +207,26 @@ export class Task { this.attemptApiRequest(previousApiReqIndex, lastApiReqIndex, shouldCompact), recursivelyMakeDiracRequests: (userContent, includeFileDetails) => this.recursivelyMakeDiracRequests(userContent, includeFileDetails), - handleEmptyAssistantResponse: (params) => this.handleEmptyAssistantResponse(params), + handleEmptyAssistantResponse: (params) => this.responseProcessor.handleEmptyAssistantResponse(params), + } + } + + private get requestLoopContext(): TaskRequestLoopContext { + return { + ...this.requestBuilderContext, + ...this.requestOutcomeContext, + steeringContext: this.steeringContext, + handleMistakeLimitReached: (userContent) => this.handleMistakeLimitReached(userContent), + enqueuePreRequestSteeringMessages: () => this.enqueuePreRequestSteeringMessages(), + resetStreamingState: () => this.resetStreamingState(), + initializeCheckpoints: (isFirstRequest) => this.initializeCheckpoints(isFirstRequest), + determineContextCompaction: (previousApiReqIndex) => this.determineContextCompaction(previousApiReqIndex), + localConversationCompaction: this.localConversationCompaction, + responseProcessor: this.responseProcessor, + streamHandler: this.streamHandler, + modelContextTracker: this.modelContextTracker, + diffViewProvider: this.diffViewProvider, + ulid: this.ulid, } } @@ -1035,14 +1051,6 @@ export class Task { return this.apiConversationManager.handleContextWindowExceededError() } - /** - * Build the system prompt, tool snapshot, and context management metadata - * for an API request. Extracted from attemptApiRequest for single-responsibility. - */ - private async buildApiRequestParams(params: { previousApiReqIndex: number; shouldCompact?: boolean }) { - return buildApiRequestParams(this.requestBuilderContext, params) - } - private async handleApiRequestError(params: { error: unknown previousApiReqIndex: number @@ -1069,113 +1077,8 @@ export class Task { this.taskState.activeVoiceStreamId = undefined } - private async processStreamResult(params: { - assistantHasContent: boolean - stopReason?: string - userContent: DiracContent[] - metricsManager: StreamingMetricsManager - modelInfo: DiracMessageModelInfo - providerId: string - model: { id: string } - }): Promise { - return processStreamResult(this.requestOutcomeContext, params) - } - async *attemptApiRequest(previousApiReqIndex: number, lastApiReqIndex: number, shouldCompact?: boolean): ApiStream { - const { systemPrompt, toolSnapshot, contextManagementMetadata, providerInfo } = await this.buildApiRequestParams({ - previousApiReqIndex, - shouldCompact, - }) - const { model, providerId } = providerInfo - - const metricsManager = new StreamingMetricsManager(this.messageStateHandler, lastApiReqIndex, this.api) - - const finalizeApiReqMsg = async (cancelReason?: DiracApiReqCancelReason, streamingFailedMessage?: string) => { - await metricsManager.updateApiReqMsgFromMetrics(cancelReason, streamingFailedMessage) - await this.messageStateHandler.updateDiracMessage(lastApiReqIndex, {}) - this.taskState.isApiRequestActive = false - this.taskState.activeVoiceStreamId = undefined - } - - const abortStream = async (cancelReason: DiracApiReqCancelReason, streamingFailedMessage?: string) => { - this.taskState.didFinishAbortingStream = true - await finalizeApiReqMsg(cancelReason, streamingFailedMessage) - this.taskState.isApiRequestActive = false - this.taskState.activeVoiceStreamId = undefined - } - - await this.appendQueuedSteeringToNextApiRequest(contextManagementMetadata.truncatedConversationHistory) - - const providerDispatch = await this.apiConversationManager.prepareProviderConversationDispatch({ - systemPrompt, - tools: toolSnapshot.nativeTools, - truncatedMessages: contextManagementMetadata.truncatedConversationHistory as DiracStorageMessage[], - providerId, - modelId: model.id, - }) - - if (this.taskState.abort) throw new Error("Task instance aborted") - - const stream = this.api.createMessage( - systemPrompt, - providerDispatch.messages.map(removeProviderBoundaryMetadataFromMessage), - toolSnapshot.nativeTools, - providerDispatch.options, - ) - const iterator = stream[Symbol.asyncIterator]() - - try { - this.taskState.status = TaskStatus.WAITING_FOR_API - - this.taskState.isWaitingForFirstChunk = true - const firstChunk = await iterator.next() - this.taskState.isWaitingForFirstChunk = false - - if (firstChunk.done) { - await finalizeApiReqMsg() - return - } - - yield firstChunk.value - - for await (const chunk of iterator) { - if (this.taskState.abort) { - await abortStream("user_cancelled") - return - } - - if (chunk.type === "usage") { - metricsManager.updateFromChunk(chunk) - yield chunk - continue - } - - yield chunk - } - - recordSuccessfulModelProviderPreset( - this.stateManager, - providerId as import("@shared/api").ApiProvider, - model.id, - model.info, - providerInfo.mode, - ) - await finalizeApiReqMsg() - } catch (error) { - const shouldRetry = await this.handleApiRequestError({ - error, - previousApiReqIndex, - lastApiReqIndex, - shouldCompact, - model, - providerId, - metricsManager, - }) - if (shouldRetry) { - yield* this.attemptApiRequest(previousApiReqIndex, lastApiReqIndex, shouldCompact) - } - return - } + yield* attemptApiRequest(this.requestLoopContext, previousApiReqIndex, lastApiReqIndex, shouldCompact) } async presentAssistantMessage() { @@ -1183,395 +1086,9 @@ export class Task { } async recursivelyMakeDiracRequests(userContent: DiracContent[], includeFileDetails = false): Promise { - this.taskState.status = TaskStatus.PREPARING - - if (this.taskState.abort) { - throw new Error("Task instance aborted") - } - await this.enqueuePreRequestSteeringMessages() - - const { model, providerId, customPrompt, mode } = this.getCurrentProviderInfo() - if (providerId && model.id) { - try { - await this.modelContextTracker.recordModelUsage(providerId, model.id, mode) - } catch (error) { - Logger.error("Failed to record model usage:", error) - } - } - - const modelInfo: DiracMessageModelInfo = { - modelId: model.id, - providerId: providerId, - mode: mode, - } - - const mistakeResult = await this.handleMistakeLimitReached(userContent) - if (mistakeResult.didEndLoop) { - return true - } - userContent = mistakeResult.userContent - - const previousApiReqIndex = findLastIndex( - this.messageStateHandler.getDiracMessages(), - (m) => m.content.type === DiracMessageType.API_STATUS, - ) - const isFirstRequest = - this.messageStateHandler.getDiracMessages().filter((m) => m.content.type === DiracMessageType.API_STATUS).length === 0 - - await this.initializeCheckpoints(isFirstRequest) - - const useCompactPrompt = customPrompt === "compact" && isLocalModel(this.getCurrentProviderInfo()) - let shouldCompact = await this.determineContextCompaction(previousApiReqIndex) - if (shouldCompact && this.localConversationCompaction.isAvailable()) { - const continuation = await this.localConversationCompaction.run({ - source: "automatic", - triggerApiRequestIndex: previousApiReqIndex, - }) - if (!continuation) return true - - const compactedContext: DiracContent[] = [{ type: "text", text: continuation }] - if (this.taskState.pinnedContext) { - compactedContext.push({ type: "text", text: this.taskState.pinnedContext }) - } - userContent = [...compactedContext, ...userContent] - shouldCompact = false - } - const steeringClaim = await this.appendQueuedSteeringToUserContent(userContent) - - this.taskState.status = TaskStatus.BUILDING_REQUEST - - let apiRequestData: Awaited> - let steeringClaimConsumed = false - try { - apiRequestData = await this.apiConversationManager.prepareApiRequest({ - userContent, - shouldCompact, - includeFileDetails, - useCompactPrompt, - previousApiReqIndex, - isFirstRequest, - providerId, - modelId: model.id, - mode: modelInfo.mode, - afterUserContentPersisted: async () => { - steeringClaimConsumed = true - if (!steeringClaim) return - await settleConsumedSteeringClaim(this.steeringContext, steeringClaim) - }, - }) - if (steeringClaim && !steeringClaimConsumed) { - if (apiRequestData.didConsumeUserContent) { - steeringClaimConsumed = true - await settleConsumedSteeringClaim(this.steeringContext, steeringClaim) - } else { - await rollbackSteeringClaim(this.steeringContext, steeringClaim.id) - } - } - } catch (error) { - if (steeringClaim && !steeringClaimConsumed) await rollbackSteeringClaim(this.steeringContext, steeringClaim.id) - throw error - } - userContent = apiRequestData.userContent - const lastApiReqIndex = apiRequestData.lastApiReqIndex - - if (apiRequestData.isDirectResponse) { - if (apiRequestData.directResponseText) { - await this.taskMessenger.upsertText(apiRequestData.directResponseText) - } - return true - } - - try { - const metricsManager = new StreamingMetricsManager(this.messageStateHandler, lastApiReqIndex, this.api) - let didFinalizeApiReqMsg = false - let usageChunkSideEffectsQueue = Promise.resolve() - - const queueUsageChunkSideEffects = ( - usageInputTokens: number, - usageOutputTokens: number, - chunkOptions?: { cacheWriteTokens?: number; cacheReadTokens?: number; totalCost?: number; stopReason?: string }, - ) => { - usageChunkSideEffectsQueue = usageChunkSideEffectsQueue.then(async () => { - if (didFinalizeApiReqMsg || this.taskState.abort) { - return - } - - await metricsManager.updateApiReqMsgFromMetrics() - await this.postStateToWebview() - await telemetryService.captureTokenUsage( - this.ulid, - usageInputTokens, - usageOutputTokens, - providerId, - model.id, - chunkOptions, - ) - }) - } - - const finalizeApiReqMsg = async (cancelReason?: DiracApiReqCancelReason, streamingFailedMessage?: string) => { - didFinalizeApiReqMsg = true - await usageChunkSideEffectsQueue - await metricsManager.updateApiReqMsgFromMetrics(cancelReason, streamingFailedMessage) - - const metrics = metricsManager.getMetrics() - this.taskState.totalInputTokens += metrics.inputTokens - this.taskState.totalOutputTokens += metrics.outputTokens - this.taskState.totalReasoningTokens += metrics.reasoningTokens - this.taskState.totalCacheWriteTokens += metrics.cacheWriteTokens - this.taskState.totalCacheReadTokens += metrics.cacheReadTokens - const cost = metricsManager.getTotalCost() - if (cost !== undefined) this.taskState.totalCost += cost - - const currentApiReqIndex = findLastIndex( - this.messageStateHandler.getDiracMessages(), - (m) => m.content.type === DiracMessageType.API_STATUS, - ) - if (currentApiReqIndex !== -1) { - this.taskState.isApiRequestActive = false - this.taskState.activeVoiceStreamId = undefined - } - } - - const abortStream = async (cancelReason: DiracApiReqCancelReason, streamingFailedMessage?: string) => { - Session.get().finalizeRequest() - - if (this.diffViewProvider.isEditing) { - await this.diffViewProvider.revertChanges() - } - - this.taskState.isApiRequestActive = false - this.taskState.activeVoiceStreamId = undefined - await finalizeApiReqMsg(cancelReason, streamingFailedMessage) - await this.messageStateHandler.saveDiracMessagesAndUpdateHistory() - - const metrics = metricsManager.getMetrics() - await this.messageStateHandler.addToApiConversationHistory({ - role: "assistant", - content: [ - { - type: "text", - text: - assistantMessage + - `\n\n[${ - cancelReason === "streaming_failed" - ? "Response interrupted by API Error" - : "Response interrupted by user" - }]`, - }, - ], - modelInfo, - metrics: { - tokens: { - prompt: metrics.inputTokens, - completion: metrics.outputTokens, - cached: (metrics.cacheWriteTokens ?? 0) + (metrics.cacheReadTokens ?? 0), - }, - cost: metrics.totalCost, - }, - ts: Date.now(), - }) - - telemetryService.captureConversationTurnEvent( - this.ulid, - providerId, - modelInfo.modelId, - "assistant", - modelInfo.mode, - undefined, - this.taskState.useNativeToolCalls, - ) - - this.taskState.didFinishAbortingStream = true - } - - await this.resetStreamingState() - - const { toolUseHandler, reasonsHandler } = this.streamHandler.getHandlers() - const stream = this.attemptApiRequest(previousApiReqIndex, lastApiReqIndex, shouldCompact) - - let assistantMessageId = "" - let assistantMessage = "" - let assistantTextOnly = "" - let assistantTextSignature: string | undefined - - let didReceiveUsageChunk = false - let stopReason: string | undefined - let didFinalizeReasoningForUi = false - - const finalizePendingReasoningMessage = async (thinking: string): Promise => { - const activeVoiceStreamId = this.taskState.activeVoiceStreamId - if (!activeVoiceStreamId) { - return false - } - - const messages = this.messageStateHandler.getDiracMessages() - const pendingReasoningIndex = messages.findIndex((m) => m.id === activeVoiceStreamId) - - if (pendingReasoningIndex !== -1) { - const msg = messages[pendingReasoningIndex] - if (msg.content.type === DiracMessageType.MARKDOWN && msg.content.isReasoning) { - await this.messageStateHandler.updateDiracMessage(pendingReasoningIndex, { - content: { type: DiracMessageType.MARKDOWN, content: thinking, isReasoning: true }, - }) - const completedReasoning = this.messageStateHandler.getDiracMessages()[pendingReasoningIndex] - if (completedReasoning) { - await this.postStateToWebview() - } - this.taskState.activeVoiceStreamId = undefined - return true - } - } - return false - } - - Session.get().startApiCall() - this.taskState.isApiRequestActive = true - let streamCoordinator: StreamChunkCoordinator | undefined - - try { - streamCoordinator = new StreamChunkCoordinator(stream, { - onUsageChunk: (chunk) => { - this.streamHandler.setRequestId(chunk.id) - didReceiveUsageChunk = true - metricsManager.updateFromChunk(chunk) - stopReason = chunk.stopReason ?? stopReason - queueUsageChunkSideEffects(chunk.inputTokens, chunk.outputTokens, { - cacheWriteTokens: chunk.cacheWriteTokens, - cacheReadTokens: chunk.cacheReadTokens, - totalCost: chunk.totalCost, - stopReason: chunk.stopReason, - }) - }, - }) - - const streamResult = await this.responseProcessor.consumeStream(streamCoordinator, { - abortStream, - finalizePendingReasoningMessage, - apiAbort: () => this.api.abort?.(), - }) - - assistantMessage = streamResult.assistantMessage - assistantTextOnly = streamResult.assistantTextOnly - assistantTextSignature = streamResult.assistantTextSignature - assistantMessageId = streamResult.assistantMessageId - didFinalizeReasoningForUi = streamResult.didFinalizeReasoningForUi - const shouldInterruptStream = streamResult.shouldInterruptStream - - if (shouldInterruptStream) { - await streamCoordinator.stop() - } else { - await streamCoordinator.waitForCompletion() - } - await usageChunkSideEffectsQueue - - if (!this.taskState.abort && !didFinalizeReasoningForUi) { - const finalReasoning = reasonsHandler.getCurrentReasoning() - if (finalReasoning?.thinking) { - await finalizePendingReasoningMessage(finalReasoning.thinking) - didFinalizeReasoningForUi = true - } - } - } catch (error) { - await streamCoordinator?.stop() - if (this.taskState.abort || this.taskState.abandoned) { - return true - } - - const diracError = ErrorService.get().toDiracError(error, this.api.getModel().id) - const errorMessage = diracError.serialize() - await this.abortTask() - await abortStream("streaming_failed", errorMessage) - await this.reinitExistingTaskFromId(this.taskId) - return true - } finally { - Session.get().endApiCall() - } - - if (!didReceiveUsageChunk) { - const apiStreamUsage = await this.api.getApiStreamUsage?.() - if (apiStreamUsage) { - metricsManager.updateFromChunk(apiStreamUsage) - queueUsageChunkSideEffects(apiStreamUsage.inputTokens, apiStreamUsage.outputTokens, { - cacheWriteTokens: apiStreamUsage.cacheWriteTokens, - cacheReadTokens: apiStreamUsage.cacheReadTokens, - totalCost: apiStreamUsage.totalCost, - stopReason: apiStreamUsage.stopReason, - }) - } - } - - const autoRetryApiReqIndex = findLastIndex( - this.messageStateHandler.getDiracMessages(), - (m) => m.content.type === DiracMessageType.API_STATUS, - ) - if (autoRetryApiReqIndex !== -1) { - const diracMessages = this.messageStateHandler.getDiracMessages() - const msg = diracMessages[autoRetryApiReqIndex] - if (msg.content.type === DiracMessageType.API_STATUS) { - const content = msg.content as Extract - const currentApiReqInfo = { ...content.status } - delete currentApiReqInfo.retryStatus - await this.messageStateHandler.updateDiracMessage(autoRetryApiReqIndex, { - content: { - type: DiracMessageType.API_STATUS, - status: currentApiReqInfo, - }, - }) - } - } - - await finalizeApiReqMsg() - await this.persistApiStopReason(stopReason) - await this.messageStateHandler.saveDiracMessagesAndUpdateHistory() - await this.postStateToWebview() - - if (this.taskState.abort) { - throw new Error("Dirac instance aborted") - } - - const assistantHasContent = await this.processAssistantResponse({ - assistantMessage, - assistantTextOnly, - assistantTextSignature, - assistantMessageId, - providerId, - modelId: model.id, - mode: modelInfo.mode, - taskMetrics: metricsManager.getMetrics(), - modelInfo, - toolUseHandler, - }) - - return await this.processStreamResult({ - assistantHasContent, - stopReason, - userContent, - metricsManager, - modelInfo, - providerId, - model, - }) - } catch (error) { - if (this.taskState.abort) { - // User-initiated abort — not a fatal error, no card needed - return true - } - const diracError = ErrorService.get().toDiracError(error) - Logger.error("[Task] Fatal error in task loop:", diracError.serialize()) - try { - const card = await this.taskMessenger.createCard({ - status: CardStatus.ERROR, - header: "Task Error", - body: `The task encountered an unexpected error and had to stop.\n\n${diracError.toDisplayMessage()}`, - }) - await card.finalize(CardStatus.ERROR) - } catch (sayError) { - Logger.error("[Task] Failed to emit error message:", sayError) - } - return true - } + return recursivelyMakeDiracRequests(this.requestLoopContext, userContent, includeFileDetails) } + private async initializeCheckpoints(isFirstRequest: boolean): Promise { return this.lifecycleManager.initializeCheckpoints(isFirstRequest) } @@ -1580,39 +1097,4 @@ export class Task { return this.apiConversationManager.determineContextCompaction(previousApiReqIndex) } - private async processAssistantResponse(params: { - assistantMessage: string - assistantTextOnly: string - assistantTextSignature?: string - assistantMessageId: string - providerId: string - modelId: string - mode: Mode - taskMetrics: { - inputTokens: number - outputTokens: number - cacheWriteTokens: number - cacheReadTokens: number - totalCost?: number - } - modelInfo: DiracMessageModelInfo - toolUseHandler: ReturnType["toolUseHandler"] - }): Promise { - return this.responseProcessor.routeAssistantResponse(params) - } - - private async handleEmptyAssistantResponse(params: { - modelInfo: DiracMessageModelInfo - taskMetrics: { - inputTokens: number - outputTokens: number - cacheWriteTokens: number - cacheReadTokens: number - totalCost?: number - } - providerId: string - model: any - }): Promise { - return this.responseProcessor.handleEmptyAssistantResponse(params) - } } From a32675a4c820f5cd1993b3ae13a370ed4c711508 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 9 Aug 2026 15:30:46 +0000 Subject: [PATCH 8/8] FB-15a: split attemptApiRequest into TaskApiRequestAttempt.ts Co-Authored-By: Alexandros Salapatas --- src/core/task/TaskApiRequestAttempt.ts | 114 +++++++++++++++++++++++++ src/core/task/TaskRequestLoop.ts | 114 +------------------------ src/core/task/index.ts | 3 +- 3 files changed, 120 insertions(+), 111 deletions(-) create mode 100644 src/core/task/TaskApiRequestAttempt.ts diff --git a/src/core/task/TaskApiRequestAttempt.ts b/src/core/task/TaskApiRequestAttempt.ts new file mode 100644 index 00000000..0cdbb481 --- /dev/null +++ b/src/core/task/TaskApiRequestAttempt.ts @@ -0,0 +1,114 @@ +import { ApiStream } from "@core/api/transform/stream" +import { recordSuccessfulModelProviderPreset } from "@core/models/modelProviderPresets" +import type { ApiProvider } from "@shared/api" +import type { DiracApiReqCancelReason } from "@shared/ExtensionMessage" +import { TaskStatus } from "@shared/ExtensionMessage" +import { removeProviderBoundaryMetadataFromMessage } from "@shared/messages/content" +import type { DiracStorageMessage } from "@shared/messages/content" +import { StreamingMetricsManager } from "./StreamingMetricsManager" +import { buildApiRequestParams } from "./TaskRequestBuilder" +import { handleApiRequestError } from "./TaskRequestOutcome" +import { appendQueuedSteeringToNextApiRequest } from "./TaskSteering" +import type { TaskRequestLoopContext } from "./TaskRequestLoop" + +export async function* attemptApiRequest( + ctx: TaskRequestLoopContext, + previousApiReqIndex: number, + lastApiReqIndex: number, + shouldCompact?: boolean, +): ApiStream { + const { systemPrompt, toolSnapshot, contextManagementMetadata, providerInfo } = await buildApiRequestParams(ctx, { + previousApiReqIndex, + shouldCompact, + }) + const { model, providerId } = providerInfo + + const metricsManager = new StreamingMetricsManager(ctx.messageStateHandler, lastApiReqIndex, ctx.api) + + const finalizeApiReqMsg = async (cancelReason?: DiracApiReqCancelReason, streamingFailedMessage?: string) => { + await metricsManager.updateApiReqMsgFromMetrics(cancelReason, streamingFailedMessage) + await ctx.messageStateHandler.updateDiracMessage(lastApiReqIndex, {}) + ctx.taskState.isApiRequestActive = false + ctx.taskState.activeVoiceStreamId = undefined + } + + const abortStream = async (cancelReason: DiracApiReqCancelReason, streamingFailedMessage?: string) => { + ctx.taskState.didFinishAbortingStream = true + await finalizeApiReqMsg(cancelReason, streamingFailedMessage) + ctx.taskState.isApiRequestActive = false + ctx.taskState.activeVoiceStreamId = undefined + } + + await appendQueuedSteeringToNextApiRequest(ctx.steeringContext, contextManagementMetadata.truncatedConversationHistory) + + const providerDispatch = await ctx.apiConversationManager.prepareProviderConversationDispatch({ + systemPrompt, + tools: toolSnapshot.nativeTools, + truncatedMessages: contextManagementMetadata.truncatedConversationHistory as DiracStorageMessage[], + providerId, + modelId: model.id, + }) + + if (ctx.taskState.abort) throw new Error("Task instance aborted") + + const stream = ctx.api.createMessage( + systemPrompt, + providerDispatch.messages.map(removeProviderBoundaryMetadataFromMessage), + toolSnapshot.nativeTools, + providerDispatch.options, + ) + const iterator = stream[Symbol.asyncIterator]() + + try { + ctx.taskState.status = TaskStatus.WAITING_FOR_API + + ctx.taskState.isWaitingForFirstChunk = true + const firstChunk = await iterator.next() + ctx.taskState.isWaitingForFirstChunk = false + + if (firstChunk.done) { + await finalizeApiReqMsg() + return + } + + yield firstChunk.value + + for await (const chunk of iterator) { + if (ctx.taskState.abort) { + await abortStream("user_cancelled") + return + } + + if (chunk.type === "usage") { + metricsManager.updateFromChunk(chunk) + yield chunk + continue + } + + yield chunk + } + + recordSuccessfulModelProviderPreset( + ctx.stateManager, + providerId as ApiProvider, + model.id, + model.info, + providerInfo.mode, + ) + await finalizeApiReqMsg() + } catch (error) { + const shouldRetry = await handleApiRequestError(ctx, { + error, + previousApiReqIndex, + lastApiReqIndex, + shouldCompact, + model, + providerId, + metricsManager, + }) + if (shouldRetry) { + yield* attemptApiRequest(ctx, previousApiReqIndex, lastApiReqIndex, shouldCompact) + } + return + } +} diff --git a/src/core/task/TaskRequestLoop.ts b/src/core/task/TaskRequestLoop.ts index 61d3775a..8b6cf63b 100644 --- a/src/core/task/TaskRequestLoop.ts +++ b/src/core/task/TaskRequestLoop.ts @@ -1,8 +1,6 @@ -import { ApiStream } from "@core/api/transform/stream" -import { recordSuccessfulModelProviderPreset } from "@core/models/modelProviderPresets" import { ErrorService } from "@services/error" import { telemetryService } from "@services/telemetry" -import type { ApiProvider } from "@shared/api" + import { findLastIndex } from "@shared/array" import type { DiracApiReqCancelReason, DiracMessageContent } from "@shared/ExtensionMessage" import { CardStatus, DiracMessageType, TaskStatus } from "@shared/ExtensionMessage" @@ -17,19 +15,17 @@ import { StreamChunkCoordinator } from "./StreamChunkCoordinator" import { StreamingMetricsManager } from "./StreamingMetricsManager" import type { StreamResponseHandler } from "./StreamResponseHandler" -import { removeProviderBoundaryMetadataFromMessage } from "@shared/messages/content" -import type { DiracContent, DiracStorageMessage } from "@shared/messages/content" +import type { DiracContent } from "@shared/messages/content" import type { DiracMessageModelInfo } from "@shared/messages/metrics" import { Logger } from "@shared/services/Logger" -import { buildApiRequestParams, type TaskRequestBuilderContext } from "./TaskRequestBuilder" +import { type TaskRequestBuilderContext } from "./TaskRequestBuilder" import { - handleApiRequestError, persistApiStopReason, processStreamResult, type TaskRequestOutcomeContext, } from "./TaskRequestOutcome" +import { attemptApiRequest } from "./TaskApiRequestAttempt" import { - appendQueuedSteeringToNextApiRequest, appendQueuedSteeringToUserContent, rollbackSteeringClaim, settleConsumedSteeringClaim, @@ -52,108 +48,6 @@ export interface TaskRequestLoopContext extends TaskRequestBuilderContext, TaskR ulid: string } -export async function* attemptApiRequest( - ctx: TaskRequestLoopContext, - previousApiReqIndex: number, - lastApiReqIndex: number, - shouldCompact?: boolean, -): ApiStream { - const { systemPrompt, toolSnapshot, contextManagementMetadata, providerInfo } = await buildApiRequestParams(ctx, { - previousApiReqIndex, - shouldCompact, - }) - const { model, providerId } = providerInfo - - const metricsManager = new StreamingMetricsManager(ctx.messageStateHandler, lastApiReqIndex, ctx.api) - - const finalizeApiReqMsg = async (cancelReason?: DiracApiReqCancelReason, streamingFailedMessage?: string) => { - await metricsManager.updateApiReqMsgFromMetrics(cancelReason, streamingFailedMessage) - await ctx.messageStateHandler.updateDiracMessage(lastApiReqIndex, {}) - ctx.taskState.isApiRequestActive = false - ctx.taskState.activeVoiceStreamId = undefined - } - - const abortStream = async (cancelReason: DiracApiReqCancelReason, streamingFailedMessage?: string) => { - ctx.taskState.didFinishAbortingStream = true - await finalizeApiReqMsg(cancelReason, streamingFailedMessage) - ctx.taskState.isApiRequestActive = false - ctx.taskState.activeVoiceStreamId = undefined - } - - await appendQueuedSteeringToNextApiRequest(ctx.steeringContext, contextManagementMetadata.truncatedConversationHistory) - - const providerDispatch = await ctx.apiConversationManager.prepareProviderConversationDispatch({ - systemPrompt, - tools: toolSnapshot.nativeTools, - truncatedMessages: contextManagementMetadata.truncatedConversationHistory as DiracStorageMessage[], - providerId, - modelId: model.id, - }) - - if (ctx.taskState.abort) throw new Error("Task instance aborted") - - const stream = ctx.api.createMessage( - systemPrompt, - providerDispatch.messages.map(removeProviderBoundaryMetadataFromMessage), - toolSnapshot.nativeTools, - providerDispatch.options, - ) - const iterator = stream[Symbol.asyncIterator]() - - try { - ctx.taskState.status = TaskStatus.WAITING_FOR_API - - ctx.taskState.isWaitingForFirstChunk = true - const firstChunk = await iterator.next() - ctx.taskState.isWaitingForFirstChunk = false - - if (firstChunk.done) { - await finalizeApiReqMsg() - return - } - - yield firstChunk.value - - for await (const chunk of iterator) { - if (ctx.taskState.abort) { - await abortStream("user_cancelled") - return - } - - if (chunk.type === "usage") { - metricsManager.updateFromChunk(chunk) - yield chunk - continue - } - - yield chunk - } - - recordSuccessfulModelProviderPreset( - ctx.stateManager, - providerId as ApiProvider, - model.id, - model.info, - providerInfo.mode, - ) - await finalizeApiReqMsg() - } catch (error) { - const shouldRetry = await handleApiRequestError(ctx, { - error, - previousApiReqIndex, - lastApiReqIndex, - shouldCompact, - model, - providerId, - metricsManager, - }) - if (shouldRetry) { - yield* attemptApiRequest(ctx, previousApiReqIndex, lastApiReqIndex, shouldCompact) - } - return - } -} - export async function recursivelyMakeDiracRequests( ctx: TaskRequestLoopContext, userContent: DiracContent[], diff --git a/src/core/task/index.ts b/src/core/task/index.ts index c131c9ca..6b30a119 100644 --- a/src/core/task/index.ts +++ b/src/core/task/index.ts @@ -81,7 +81,8 @@ import { persistApiStopReason, type TaskRequestOutcomeContext, } from "./TaskRequestOutcome" -import { attemptApiRequest, recursivelyMakeDiracRequests, type TaskRequestLoopContext } from "./TaskRequestLoop" +import { attemptApiRequest } from "./TaskApiRequestAttempt" +import { recursivelyMakeDiracRequests, type TaskRequestLoopContext } from "./TaskRequestLoop" import { TaskState } from "./TaskState" import { appendQueuedSteeringToNextApiRequest,