Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions .pi/extensions/workflow-orchestrator/config.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
Expand Down Expand Up @@ -54,6 +66,7 @@ const WorkflowSchema = Type.Object({
waveSource: WaveSourceSchema,
taskFlow: Type.Object({
stages: Type.Array(StageSchema),
memory: Type.Optional(TaskFlowMemorySchema),
}),
});

Expand Down Expand Up @@ -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");
Expand Down
18 changes: 12 additions & 6 deletions .pi/extensions/workflow-orchestrator/engine.ts
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ export interface TaskFlowInput<TTask, TOutput> {
stageId: string,
output: TOutput | null,
error?: string,
reason?: "verification_failed" | "malformed_output" | "error",
) => boolean;
applyGenericFailure: (task: TTask, error: string) => boolean;
}
Expand Down Expand Up @@ -73,7 +74,7 @@ export async function runTaskFlow<TTask extends { retries: number }, TOutput>(
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;
Expand All @@ -93,18 +94,22 @@ export async function runTaskFlow<TTask extends { retries: number }, TOutput>(
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);
Expand All @@ -115,9 +120,10 @@ export async function runTaskFlow<TTask extends { retries: number }, TOutput>(
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;
Expand Down
41 changes: 40 additions & 1 deletion .pi/extensions/workflow-orchestrator/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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] : []);
Expand All @@ -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);
},
Expand Down
1 change: 1 addition & 0 deletions .pi/extensions/workflow-orchestrator/state.ts
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ export interface TaskState extends WorkflowTask {
lastOutput?: string;
lastActivityAt?: number;
sessionFiles?: Record<string, string>;
sessionResetCounts?: Record<string, number>;
resumeMessage?: string;
}

Expand Down
5 changes: 5 additions & 0 deletions .pi/workflows/default.workflow.json
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,11 @@
"type": "pm"
},
"taskFlow": {
"memory": {
"keepDeveloperMemory": true,
"keepVerifierMemoryOnDeveloperFailure": true,
"verifierSelfFailureMemory": "keep"
},
"stages": [
{
"id": "develop",
Expand Down
6 changes: 6 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Expand Down
41 changes: 41 additions & 0 deletions tests/config-extended.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {
Expand Down
25 changes: 25 additions & 0 deletions tests/config.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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", () => {
Expand Down
50 changes: 48 additions & 2 deletions tests/engine.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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");
});
});
Loading