diff --git a/.pi/extensions/workflow-orchestrator/config.ts b/.pi/extensions/workflow-orchestrator/config.ts index 0e143ac..b0d427d 100644 --- a/.pi/extensions/workflow-orchestrator/config.ts +++ b/.pi/extensions/workflow-orchestrator/config.ts @@ -19,6 +19,18 @@ const StageSchema = Type.Object({ transitions: Type.Optional(Type.Array(TransitionSchema)), }); +const TaskFlowMemorySchema = Type.Object({ + keepDeveloperMemory: Type.Optional(Type.Boolean()), + keepVerifierMemoryOnDeveloperFailure: Type.Optional(Type.Boolean()), + verifierSelfFailureMemory: Type.Optional( + Type.Union([ + Type.Literal("keep"), + Type.Literal("reset"), + Type.Literal("reset_on_malformed_output"), + ]), + ), +}); + const TaskSchema = Type.Object({ id: Type.String(), title: Type.String(), @@ -54,6 +66,7 @@ const WorkflowSchema = Type.Object({ waveSource: WaveSourceSchema, taskFlow: Type.Object({ stages: Type.Array(StageSchema), + memory: Type.Optional(TaskFlowMemorySchema), }), }); @@ -114,6 +127,12 @@ export function loadWorkflowConfig(cwd: string, name: string): LoadedWorkflow { config.maxWaves = config.maxWaves ?? 10; config.maxTaskRetries = config.maxTaskRetries ?? 2; config.parallelism = config.parallelism ?? 1; + config.taskFlow.memory = config.taskFlow.memory ?? {}; + config.taskFlow.memory.keepDeveloperMemory = config.taskFlow.memory.keepDeveloperMemory ?? true; + config.taskFlow.memory.keepVerifierMemoryOnDeveloperFailure = + config.taskFlow.memory.keepVerifierMemoryOnDeveloperFailure ?? true; + config.taskFlow.memory.verifierSelfFailureMemory = + config.taskFlow.memory.verifierSelfFailureMemory ?? "keep"; if (config.parallelism < 1) { throw new Error("parallelism must be at least 1"); diff --git a/.pi/extensions/workflow-orchestrator/engine.ts b/.pi/extensions/workflow-orchestrator/engine.ts index 7768f8e..a494620 100644 --- a/.pi/extensions/workflow-orchestrator/engine.ts +++ b/.pi/extensions/workflow-orchestrator/engine.ts @@ -28,6 +28,7 @@ export interface TaskFlowInput { stageId: string, output: TOutput | null, error?: string, + reason?: "verification_failed" | "malformed_output" | "error", ) => boolean; applyGenericFailure: (task: TTask, error: string) => boolean; } @@ -73,7 +74,7 @@ export async function runTaskFlow( input.onError?.(stage, input.task, error instanceof Error ? error : new Error(message)); if (stage.id === "verify") { - const retry = input.applyVerifyFailure(input.task, stage.id, null, message); + const retry = input.applyVerifyFailure(input.task, stage.id, null, message, "error"); if (!retry) { input.markFailed(input.task, stage.id); return; @@ -93,18 +94,22 @@ export async function runTaskFlow( if (input.isStopped?.(input.task)) return; let nextStageId: string | undefined; + const firstStageId = input.startStageId ?? input.stages[0]?.id; + let matchedTransition = false; if (stage.transitions && stage.transitions.length > 0) { const fieldTarget = (output as any)?.output ?? output; for (const transition of stage.transitions) { const fieldValue = getField(fieldTarget as any, transition.when.field); if (String(fieldValue) === transition.when.equals) { + matchedTransition = true; nextStageId = transition.next; break; } } - // If no transition matched (e.g., status is "unknown" or missing), retry verifier + // If no transition matched (e.g., status is "unknown" or missing), keep verifier + // running and let verify-failure retry handling decide whether to continue/fail. if (!nextStageId && stage.id === "verify") { - nextStageId = stage.id; // Stay in verify stage + nextStageId = stage.id; } } else { nextStageId = getNextStageId(input.stages, stage.id); @@ -115,9 +120,10 @@ export async function runTaskFlow( return; } - const firstStageId = input.startStageId ?? input.stages[0]?.id; - if (nextStageId === firstStageId && stage.id === "verify") { - const retry = input.applyVerifyFailure(input.task, stage.id, output); + if ((nextStageId === firstStageId || nextStageId === stage.id) && stage.id === "verify") { + const reason = + nextStageId === stage.id && !matchedTransition ? "malformed_output" : "verification_failed"; + const retry = input.applyVerifyFailure(input.task, stage.id, output, undefined, reason); if (!retry) { input.markFailed(input.task, stage.id); return; diff --git a/.pi/extensions/workflow-orchestrator/index.ts b/.pi/extensions/workflow-orchestrator/index.ts index 4770beb..9b06e18 100644 --- a/.pi/extensions/workflow-orchestrator/index.ts +++ b/.pi/extensions/workflow-orchestrator/index.ts @@ -289,6 +289,22 @@ function stopTask(task: TaskState) { task.lastNote = "stopped"; } +function resetStageMemory(task: TaskState, stageId: string) { + if (!currentState) return; + const key = getRunnerKey(task.id, stageId); + const runner = taskRunners.get(key); + runner?.agent.dispose(); + taskRunners.delete(key); + + const workflowDir = path.join(".pi", "workflows", "sessions", currentState.runId); + fs.mkdirSync(workflowDir, { recursive: true }); + if (!task.sessionResetCounts) task.sessionResetCounts = {}; + const nextReset = (task.sessionResetCounts[stageId] ?? 0) + 1; + task.sessionResetCounts[stageId] = nextReset; + if (!task.sessionFiles) task.sessionFiles = {}; + task.sessionFiles[stageId] = path.join(workflowDir, `${task.id}-${stageId}-r${nextReset}.jsonl`); +} + async function messageTask( pi: ExtensionAPI, ctx: ExtensionContext, @@ -507,7 +523,7 @@ async function processTask( setState(pi, ctx, { ...currentState, tasks: [...currentState.tasks] }); }, - applyVerifyFailure: (t, stageId, result, errorMessage) => { + applyVerifyFailure: (t, stageId, result, errorMessage, reason = "verification_failed") => { if (!currentState) return false; // Guard against workflow stop during execution const output = result?.output; const issues = output?.issues ?? (errorMessage ? [errorMessage] : []); @@ -517,6 +533,29 @@ async function processTask( t.retries += 1; t.lastNote = errorMessage ? `error: ${errorMessage}` : "fail"; if (errorMessage) t.lastOutput = truncateTicker(errorMessage); + + const keepDeveloperMemory = config.taskFlow.memory?.keepDeveloperMemory ?? true; + const keepVerifierMemoryOnDeveloperFailure = + config.taskFlow.memory?.keepVerifierMemoryOnDeveloperFailure ?? true; + const verifierSelfFailureMemory = config.taskFlow.memory?.verifierSelfFailureMemory ?? "keep"; + + if ((reason === "verification_failed" || reason === "error") && !keepDeveloperMemory) { + resetStageMemory(t, "develop"); + } + if (reason === "verification_failed" && !keepVerifierMemoryOnDeveloperFailure) { + resetStageMemory(t, "verify"); + } + if (reason === "malformed_output") { + if ( + verifierSelfFailureMemory === "reset" || + verifierSelfFailureMemory === "reset_on_malformed_output" + ) { + resetStageMemory(t, "verify"); + } + } else if (reason === "error" && verifierSelfFailureMemory === "reset") { + resetStageMemory(t, "verify"); + } + setState(pi, ctx, { ...currentState, tasks: [...currentState.tasks] }); return t.retries <= (config.maxTaskRetries ?? 2); }, diff --git a/.pi/extensions/workflow-orchestrator/state.ts b/.pi/extensions/workflow-orchestrator/state.ts index f02f305..0d3bc80 100644 --- a/.pi/extensions/workflow-orchestrator/state.ts +++ b/.pi/extensions/workflow-orchestrator/state.ts @@ -14,6 +14,7 @@ export interface TaskState extends WorkflowTask { lastOutput?: string; lastActivityAt?: number; sessionFiles?: Record; + sessionResetCounts?: Record; resumeMessage?: string; } diff --git a/.pi/workflows/default.workflow.json b/.pi/workflows/default.workflow.json index f70fa45..4cb1ded 100644 --- a/.pi/workflows/default.workflow.json +++ b/.pi/workflows/default.workflow.json @@ -19,6 +19,11 @@ "type": "pm" }, "taskFlow": { + "memory": { + "keepDeveloperMemory": true, + "keepVerifierMemoryOnDeveloperFailure": true, + "verifierSelfFailureMemory": "keep" + }, "stages": [ { "id": "develop", diff --git a/README.md b/README.md index af9b4e7..c2d1d2a 100644 --- a/README.md +++ b/README.md @@ -159,6 +159,12 @@ Edit `.pi/workflows/default.workflow.json` to customize: - **maxPmRetries** - Retry limit for PM wave generation (default: 3) - **allowedExtensions** - Whitelist extensions for all subagents - **allowedExtensionsByAgent** - Per-agent extension allowlists +- **taskFlow.memory.keepDeveloperMemory** - Preserve developer session memory across retries (default: `true`) +- **taskFlow.memory.keepVerifierMemoryOnDeveloperFailure** - Preserve verifier memory after it returns `fail` and flow goes back to developer (default: `true`) +- **taskFlow.memory.verifierSelfFailureMemory** - Verifier memory policy when verifier fails **without** transitioning to developer (for example malformed/unmatched output, or verifier runtime errors): + - `keep` (default) + - `reset` + - `reset_on_malformed_output` Example: diff --git a/tests/config-extended.test.ts b/tests/config-extended.test.ts index 1c2b3a2..0fd8428 100644 --- a/tests/config-extended.test.ts +++ b/tests/config-extended.test.ts @@ -245,6 +245,47 @@ describe("config.ts - additional coverage", () => { }); }); + describe("task memory configuration", () => { + it("accepts memory configuration values", () => { + const config = { + name: "test", + goal: "test", + agents: { pm: "pm", developer: "dev", verifier: "ver" }, + waveSource: { type: "static", staticWaves: [] }, + taskFlow: { + memory: { + keepDeveloperMemory: false, + keepVerifierMemoryOnDeveloperFailure: false, + verifierSelfFailureMemory: "reset", + }, + stages: [{ id: "s1", agent: "dev", inputTemplate: "t", outputSchema: {} }], + }, + }; + createWorkflowConfig("test", JSON.stringify(config)); + const loaded = loadWorkflowConfig(tempDir, "test"); + expect(loaded.config.taskFlow.memory?.keepDeveloperMemory).toBe(false); + expect(loaded.config.taskFlow.memory?.keepVerifierMemoryOnDeveloperFailure).toBe(false); + expect(loaded.config.taskFlow.memory?.verifierSelfFailureMemory).toBe("reset"); + }); + + it("rejects invalid verifierSelfFailureMemory values", () => { + const config = { + name: "test", + goal: "test", + agents: { pm: "pm", developer: "dev", verifier: "ver" }, + waveSource: { type: "static", staticWaves: [] }, + taskFlow: { + memory: { + verifierSelfFailureMemory: "invalid", + }, + stages: [{ id: "s1", agent: "dev", inputTemplate: "t", outputSchema: {} }], + }, + }; + createWorkflowConfig("test", JSON.stringify(config)); + expect(() => loadWorkflowConfig(tempDir, "test")).toThrow(); + }); + }); + describe("task validation", () => { it("accepts task with requirements", () => { const config = { diff --git a/tests/config.test.ts b/tests/config.test.ts index e580dd5..596487c 100644 --- a/tests/config.test.ts +++ b/tests/config.test.ts @@ -37,6 +37,31 @@ describe("loadWorkflowConfig", () => { const loaded = loadWorkflowConfig(cwd, name); expect(loaded.config.name).toBe("temp"); expect(loaded.config.parallelism).toBe(1); + expect(loaded.config.taskFlow.memory?.keepDeveloperMemory).toBe(true); + expect(loaded.config.taskFlow.memory?.keepVerifierMemoryOnDeveloperFailure).toBe(true); + expect(loaded.config.taskFlow.memory?.verifierSelfFailureMemory).toBe("keep"); + }); + + it("loads custom task memory policy", () => { + const config = { + name: "temp", + goal: "test", + agents: { pm: "pm", developer: "dev", verifier: "ver" }, + waveSource: { type: "static", staticWaves: [{ goal: "g", tasks: [] }] }, + taskFlow: { + memory: { + keepDeveloperMemory: false, + keepVerifierMemoryOnDeveloperFailure: false, + verifierSelfFailureMemory: "keep", + }, + stages: [{ id: "develop", agent: "dev", inputTemplate: "x", outputSchema: {} }], + }, + }; + const { cwd, name } = setupTempConfig(JSON.stringify(config)); + const loaded = loadWorkflowConfig(cwd, name); + expect(loaded.config.taskFlow.memory?.keepDeveloperMemory).toBe(false); + expect(loaded.config.taskFlow.memory?.keepVerifierMemoryOnDeveloperFailure).toBe(false); + expect(loaded.config.taskFlow.memory?.verifierSelfFailureMemory).toBe("keep"); }); it("accepts requirements and per-agent extensions", () => { diff --git a/tests/engine.test.ts b/tests/engine.test.ts index ad1c39b..dcc104b 100644 --- a/tests/engine.test.ts +++ b/tests/engine.test.ts @@ -121,18 +121,64 @@ describe("runTaskFlow", () => { it("retries verifier when status is unknown", async () => { const task: TestTask = { id: "T1", status: "pending", retries: 0 }; let verifyCount = 0; + let developCount = 0; await runTaskFlow({ ...baseOptions(task), runStage: async (stage) => { - if (stage.id === "develop") return { output: { status: "done" }, outputText: "dev" }; + if (stage.id === "develop") { + developCount += 1; + return { output: { status: "done" }, outputText: "dev" }; + } verifyCount += 1; if (verifyCount === 1) return { output: { status: "unknown" }, outputText: "verifier failed to parse" }; return { output: { status: "pass", issues: [] }, outputText: "verify" }; }, }); + expect(developCount).toBe(1); expect(verifyCount).toBe(2); - // Note: retries counter is for dev→verify loops, not verify retries + expect(task.status).toBe("verified"); + }); + + it("fails when verifier keeps returning unknown status", async () => { + const task: TestTask = { id: "T1", status: "pending", retries: 2 }; + let verifyCount = 0; + await runTaskFlow({ + ...baseOptions(task), + maxRetries: 2, + runStage: async (stage) => { + if (stage.id === "develop") return { output: { status: "done" }, outputText: "dev" }; + verifyCount += 1; + return { output: { status: "unknown" }, outputText: "verifier unknown" }; + }, + }); + expect(verifyCount).toBe(1); + expect(task.status).toBe("failed"); + }); + + it("passes verify failure reasons for fail vs malformed outputs", async () => { + const task: TestTask = { id: "T1", status: "pending", retries: 0 }; + const reasons: string[] = []; + let verifyCount = 0; + + await runTaskFlow({ + ...baseOptions(task), + applyVerifyFailure: (t, stageId, result, error, reason) => { + reasons.push(String(reason)); + return baseOptions(t).applyVerifyFailure(t, stageId, result, error); + }, + runStage: async (stage) => { + if (stage.id === "develop") return { output: { status: "done" }, outputText: "dev" }; + verifyCount += 1; + if (verifyCount === 1) return { output: { status: "unknown" }, outputText: "unknown" }; + if (verifyCount === 2) + return { output: { status: "fail", issues: ["needs fix"] }, outputText: "fail" }; + return { output: { status: "pass", issues: [] }, outputText: "pass" }; + }, + }); + + expect(reasons).toContain("malformed_output"); + expect(reasons).toContain("verification_failed"); expect(task.status).toBe("verified"); }); });