refactor(agent): run harness loop directly

This commit is contained in:
Mario Zechner
2026-05-14 21:33:34 +02:00
parent 846906e4d1
commit b7ea82105a
11 changed files with 1711 additions and 251 deletions

View File

@@ -1,11 +1,13 @@
import { getModel } from "@earendil-works/pi-ai";
import { describe, expect, it } from "vitest";
import { fauxAssistantMessage, fauxToolCall, getModel, registerFauxProvider } from "@earendil-works/pi-ai";
import { afterEach, describe, expect, it } from "vitest";
import { AgentHarness } from "../../src/harness/agent-harness.js";
import { NodeExecutionEnv } from "../../src/harness/env/nodejs.js";
import { Session } from "../../src/harness/session/session.js";
import { InMemorySessionStorage } from "../../src/harness/session/storage/memory.js";
import type { PromptTemplate, Skill } from "../../src/harness/types.js";
import type { AgentTool } from "../../src/types.js";
import type { AgentMessage, AgentTool } from "../../src/types.js";
import { calculateTool } from "../utils/calculate.js";
import { getCurrentTimeTool } from "../utils/get-current-time.js";
interface AppSkill extends Skill {
source: "project" | "user";
@@ -19,6 +21,39 @@ interface AppTool extends AgentTool {
source: "builtin" | "extension";
}
const registrations: Array<{ unregister(): void }> = [];
function textFromUserMessages(messages: Array<{ role: string; content: unknown }>): string[] {
return messages.flatMap((message) => {
if (message.role !== "user") return [];
if (typeof message.content === "string") return [message.content];
if (!Array.isArray(message.content)) return [];
return message.content.flatMap((part) => {
if (!part || typeof part !== "object" || !("type" in part) || part.type !== "text") return [];
return "text" in part && typeof part.text === "string" ? [part.text] : [];
});
});
}
function deferred(): { promise: Promise<void>; resolve: () => void } {
let resolve = () => {};
const promise = new Promise<void>((resolvePromise) => {
resolve = resolvePromise;
});
return { promise, resolve };
}
function getReasoning(options: unknown): unknown {
if (!options || typeof options !== "object" || !("reasoning" in options)) return undefined;
return options.reasoning;
}
afterEach(() => {
for (const registration of registrations.splice(0)) {
registration.unregister();
}
});
describe("AgentHarness", () => {
it("constructs directly and exposes queue modes", () => {
const session = new Session(new InMemorySessionStorage());
@@ -28,18 +63,399 @@ describe("AgentHarness", () => {
env,
session,
model: initialModel,
thinkingLevel: "high",
systemPrompt: "You are helpful.",
steeringMode: "all",
followUpMode: "all",
});
expect(harness.env).toBe(env);
expect(harness.agent.state.model).toBe(initialModel);
expect(harness.steeringMode).toBe("all");
expect(harness.followUpMode).toBe("all");
harness.steeringMode = "one-at-a-time";
harness.followUpMode = "one-at-a-time";
expect(harness.agent.steeringMode).toBe("one-at-a-time");
expect(harness.agent.followUpMode).toBe("one-at-a-time");
expect(harness.getModel()).toBe(initialModel);
expect(harness.getThinkingLevel()).toBe("high");
expect(harness.getSteeringMode()).toBe("all");
expect(harness.getFollowUpMode()).toBe("all");
harness.setSteeringMode("one-at-a-time");
harness.setFollowUpMode("one-at-a-time");
expect(harness.getSteeringMode()).toBe("one-at-a-time");
expect(harness.getFollowUpMode()).toBe("one-at-a-time");
});
it("drains one queued steering message at a time and emits queue updates", async () => {
const registration = registerFauxProvider();
registrations.push(registration);
const userCounts: number[] = [];
registration.setResponses([
(context) => {
userCounts.push(context.messages.filter((message) => message.role === "user").length);
return fauxAssistantMessage("first");
},
(context) => {
userCounts.push(context.messages.filter((message) => message.role === "user").length);
return fauxAssistantMessage("second");
},
(context) => {
userCounts.push(context.messages.filter((message) => message.role === "user").length);
return fauxAssistantMessage("third");
},
]);
const harness = new AgentHarness({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session: new Session(new InMemorySessionStorage()),
model: registration.getModel(),
steeringMode: "one-at-a-time",
});
const steerQueueLengths: number[] = [];
let queued = false;
harness.subscribe((event) => {
if (event.type === "queue_update") {
steerQueueLengths.push(event.steer.length);
}
if (event.type === "message_start" && event.message.role === "assistant" && !queued) {
queued = true;
harness.steer("one");
harness.steer("two");
}
});
await harness.prompt("hello");
expect(userCounts).toEqual([1, 2, 3]);
expect(steerQueueLengths).toEqual([1, 2, 1, 0]);
});
it("appends before_agent_start messages and persists them", async () => {
const registration = registerFauxProvider();
registrations.push(registration);
let requestText: string[] = [];
registration.setResponses([
(context) => {
requestText = textFromUserMessages(context.messages);
return fauxAssistantMessage("ok");
},
]);
const session = new Session(new InMemorySessionStorage());
const harness = new AgentHarness({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session,
model: registration.getModel(),
});
harness.on("before_agent_start", () => ({
messages: [{ role: "user", content: [{ type: "text", text: "hook" }], timestamp: Date.now() }],
}));
await harness.prompt("hello");
const persistedText = (await session.getEntries()).flatMap((entry) => {
if (entry.type !== "message" || entry.message.role !== "user") return [];
const content = entry.message.content;
if (typeof content === "string") return [content];
return content.flatMap((part) => (part.type === "text" ? [part.text] : []));
});
expect(requestText).toEqual(["hello", "hook"]);
expect(persistedText).toEqual(["hello", "hook"]);
});
it("abort clears steer and follow-up queues but preserves next-turn messages", async () => {
const registration = registerFauxProvider();
registrations.push(registration);
let releaseFirstResponse: (() => void) | undefined;
let abortedSignal: AbortSignal | undefined;
const firstResponseReleased = new Promise<void>((resolve) => {
releaseFirstResponse = resolve;
});
const secondRequestText: string[] = [];
registration.setResponses([
async (_context, options) => {
abortedSignal = options?.signal;
await firstResponseReleased;
return fauxAssistantMessage("aborted-ish");
},
(context) => {
secondRequestText.push(...textFromUserMessages(context.messages));
return fauxAssistantMessage("second");
},
]);
const harness = new AgentHarness({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session: new Session(new InMemorySessionStorage()),
model: registration.getModel(),
});
const queueUpdates: Array<{ steer: number; followUp: number; nextTurn: number }> = [];
harness.subscribe((event) => {
if (event.type === "queue_update") {
queueUpdates.push({
steer: event.steer.length,
followUp: event.followUp.length,
nextTurn: event.nextTurn.length,
});
}
});
const firstPrompt = harness.prompt("first");
await new Promise((resolve) => setTimeout(resolve, 0));
harness.steer("steer");
harness.followUp("follow");
harness.nextTurn("next");
const abortResultPromise = harness.abort();
await new Promise((resolve) => setTimeout(resolve, 0));
expect(abortedSignal?.aborted).toBe(true);
releaseFirstResponse?.();
const abortResult = await abortResultPromise;
await firstPrompt;
await harness.prompt("second");
expect(abortResult.clearedSteer).toHaveLength(1);
expect(abortResult.clearedFollowUp).toHaveLength(1);
expect(queueUpdates).toContainEqual({ steer: 0, followUp: 0, nextTurn: 1 });
expect(secondRequestText).toEqual(["first", "next", "second"]);
});
it("drains follow-up messages one at a time after the agent would otherwise stop", async () => {
const registration = registerFauxProvider();
registrations.push(registration);
const userCounts: number[] = [];
registration.setResponses([
(context) => {
userCounts.push(context.messages.filter((message) => message.role === "user").length);
return fauxAssistantMessage("first");
},
(context) => {
userCounts.push(context.messages.filter((message) => message.role === "user").length);
return fauxAssistantMessage("second");
},
(context) => {
userCounts.push(context.messages.filter((message) => message.role === "user").length);
return fauxAssistantMessage("third");
},
]);
const harness = new AgentHarness({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session: new Session(new InMemorySessionStorage()),
model: registration.getModel(),
followUpMode: "one-at-a-time",
});
const followUpQueueLengths: number[] = [];
let queued = false;
harness.subscribe((event) => {
if (event.type === "queue_update") {
followUpQueueLengths.push(event.followUp.length);
}
if (event.type === "message_start" && event.message.role === "assistant" && !queued) {
queued = true;
harness.followUp("one");
harness.followUp("two");
}
});
await harness.prompt("hello");
expect(userCounts).toEqual([1, 2, 3]);
expect(followUpQueueLengths).toEqual([1, 2, 1, 0]);
});
it("settles thrown hook failures with persisted assistant error messages", async () => {
const registration = registerFauxProvider();
registrations.push(registration);
registration.setResponses([() => fauxAssistantMessage("should not be used")]);
const session = new Session(new InMemorySessionStorage());
const harness = new AgentHarness({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session,
model: registration.getModel(),
});
const events: string[] = [];
harness.subscribe((event) => {
events.push(event.type);
});
harness.on("context", () => {
throw new Error("context exploded");
});
const response = await harness.prompt("hello");
await expect(harness.prompt("after failure")).resolves.toMatchObject({ role: "assistant" });
const entries = await session.getEntries();
const messages = entries.flatMap((entry) => (entry.type === "message" ? [entry.message] : []));
expect(response.stopReason).toBe("error");
expect(response.errorMessage).toBe("context exploded");
expect(messages[0]?.role).toBe("user");
expect(messages[1]).toMatchObject({ role: "assistant", stopReason: "error", errorMessage: "context exploded" });
expect(events).toContain("agent_end");
expect(events).toContain("settled");
});
it("refreshes model, thinking level, resources, system prompt, and active tools at save points", async () => {
const registration = registerFauxProvider({
models: [
{ id: "first", reasoning: true },
{ id: "second", reasoning: true },
],
});
registrations.push(registration);
const secondModel = registration.getModel("second");
if (!secondModel) throw new Error("missing second faux model");
const captured: Array<{ modelId: string; reasoning: unknown; systemPrompt: string; tools: string[] }> = [];
registration.setResponses([
(context, options, _state, model) => {
captured.push({
modelId: model.id,
reasoning: getReasoning(options),
systemPrompt: context.systemPrompt ?? "",
tools: context.tools?.map((tool) => tool.name) ?? [],
});
return fauxAssistantMessage(fauxToolCall("calculate", { expression: "1 + 1" }, { id: "call-1" }), {
stopReason: "toolUse",
});
},
(context, options, _state, model) => {
captured.push({
modelId: model.id,
reasoning: getReasoning(options),
systemPrompt: context.systemPrompt ?? "",
tools: context.tools?.map((tool) => tool.name) ?? [],
});
return fauxAssistantMessage("done");
},
]);
const harness = new AgentHarness<Skill, PromptTemplate, AgentTool>({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session: new Session(new InMemorySessionStorage()),
model: registration.getModel(),
thinkingLevel: "off",
resources: {
skills: [{ name: "prompt", description: "prompt", content: "first prompt", filePath: "/skills/prompt" }],
},
systemPrompt: ({ resources }) => resources.skills?.[0]?.content ?? "missing prompt",
tools: [calculateTool],
});
harness.subscribe((event) => {
if (event.type === "tool_execution_start") {
void harness.setModel(secondModel);
void harness.setThinkingLevel("high");
void harness.setResources({
skills: [
{ name: "prompt", description: "prompt", content: "second prompt", filePath: "/skills/prompt" },
],
});
void harness.setTools([calculateTool, getCurrentTimeTool], [getCurrentTimeTool.name]);
}
});
await harness.prompt("hello");
expect(captured).toEqual([
{ modelId: "first", reasoning: undefined, systemPrompt: "first prompt", tools: ["calculate"] },
{ modelId: "second", reasoning: "high", systemPrompt: "second prompt", tools: ["get_current_time"] },
]);
});
it("orders pending listener session writes after agent-emitted messages", async () => {
const registration = registerFauxProvider();
registrations.push(registration);
registration.setResponses([() => fauxAssistantMessage("ok")]);
const session = new Session(new InMemorySessionStorage());
const harness = new AgentHarness({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session,
model: registration.getModel(),
});
let wrotePendingMessage = false;
harness.subscribe(async (event) => {
if (event.type === "message_end" && event.message.role === "assistant" && !wrotePendingMessage) {
wrotePendingMessage = true;
await harness.appendMessage({
role: "custom",
customType: "listener",
content: "listener write",
display: true,
timestamp: Date.now(),
} as AgentMessage);
}
});
await harness.prompt("hello");
const entries = await session.getEntries();
const roles = entries.flatMap((entry) => (entry.type === "message" ? [entry.message.role] : []));
expect(roles).toEqual(["user", "assistant", "custom"]);
});
it("waitForIdle waits for external run settlement and awaited listeners", async () => {
const registration = registerFauxProvider();
registrations.push(registration);
registration.setResponses([() => fauxAssistantMessage("ok")]);
const barrier = deferred();
const harness = new AgentHarness({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session: new Session(new InMemorySessionStorage()),
model: registration.getModel(),
});
let listenerFinished = false;
harness.subscribe(async (event) => {
if (event.type === "agent_end") {
await barrier.promise;
listenerFinished = true;
}
});
const promptPromise = harness.prompt("hello");
let idleResolved = false;
const idlePromise = harness.waitForIdle().then(() => {
idleResolved = true;
});
await new Promise((resolve) => setTimeout(resolve, 10));
expect(idleResolved).toBe(false);
expect(listenerFinished).toBe(false);
barrier.resolve();
await Promise.all([promptPromise, idlePromise]);
expect(idleResolved).toBe(true);
expect(listenerFinished).toBe(true);
});
it("runs tool_call and tool_result hooks through the direct loop", async () => {
const registration = registerFauxProvider();
registrations.push(registration);
registration.setResponses([
() =>
fauxAssistantMessage(fauxToolCall("calculate", { expression: "2 + 2" }, { id: "call-1" }), {
stopReason: "toolUse",
}),
]);
const session = new Session(new InMemorySessionStorage());
const harness = new AgentHarness({
env: new NodeExecutionEnv({ cwd: process.cwd() }),
session,
model: registration.getModel(),
tools: [calculateTool],
});
const seenToolCalls: Array<{ id: string; name: string; expression: unknown }> = [];
harness.on("tool_call", (event) => {
seenToolCalls.push({ id: event.toolCallId, name: event.toolName, expression: event.input.expression });
return undefined;
});
harness.on("tool_result", (event) => {
expect(event.toolCallId).toBe("call-1");
expect(event.toolName).toBe("calculate");
return {
content: [{ type: "text", text: "patched result" }],
details: { patched: true },
terminate: true,
};
});
await harness.prompt("hello");
const toolResult = (await session.getEntries()).find(
(entry) => entry.type === "message" && entry.message.role === "toolResult",
);
expect(seenToolCalls).toEqual([{ id: "call-1", name: "calculate", expression: "2 + 2" }]);
expect(toolResult).toMatchObject({
type: "message",
message: {
role: "toolResult",
content: [{ type: "text", text: "patched result" }],
details: { patched: true },
},
});
});
it("preserves app resource types for getters and update events", async () => {

View File

@@ -0,0 +1,158 @@
import { describe, expect, it } from "vitest";
import { truncateHead, truncateTail } from "../../src/harness/utils/truncate.js";
const encoder = new TextEncoder();
function byteLength(content: string): number {
return encoder.encode(content).length;
}
function bufferTail(content: string, maxBytes: number): string {
const bytes = Buffer.from(content, "utf8");
if (bytes.length <= maxBytes) return content;
let start = bytes.length - maxBytes;
while (start < bytes.length && (bytes[start] & 0xc0) === 0x80) start++;
return bytes.subarray(start).toString("utf8");
}
function assertMatchesBufferTail(input: string, maxByteValues?: readonly number[]): void {
const totalBytes = Buffer.byteLength(input, "utf8");
const values = maxByteValues ?? Array.from({ length: totalBytes + 5 }, (_, maxBytes) => maxBytes);
for (const maxBytes of values) {
const result = truncateTail(input, { maxBytes, maxLines: 10 });
const expected = bufferTail(input, maxBytes);
if (result.content !== expected) {
throw new Error(
`tail mismatch input=${JSON.stringify(input)} maxBytes=${maxBytes} expected=${JSON.stringify(expected)} actual=${JSON.stringify(result.content)}`,
);
}
const outputBytes = Buffer.byteLength(result.content, "utf8");
if (outputBytes > maxBytes) {
throw new Error(
`tail output exceeded byte limit input=${JSON.stringify(input)} maxBytes=${maxBytes} outputBytes=${outputBytes}`,
);
}
}
}
function sampledByteLimits(input: string): number[] {
const totalBytes = Buffer.byteLength(input, "utf8");
const candidates = [
0,
1,
2,
3,
4,
5,
8,
Math.floor(totalBytes / 2) - 1,
Math.floor(totalBytes / 2),
Math.floor(totalBytes / 2) + 1,
totalBytes - 8,
totalBytes - 5,
totalBytes - 4,
totalBytes - 3,
totalBytes - 2,
totalBytes - 1,
totalBytes,
totalBytes + 1,
totalBytes + 4,
];
return [...new Set(candidates.filter((value) => value >= 0))].sort((a, b) => a - b);
}
describe("truncate utilities", () => {
it("counts UTF-8 bytes without Node Buffer", () => {
const content = "aé🙂\nb";
const result = truncateHead(content, { maxBytes: 100, maxLines: 10 });
expect(result.truncated).toBe(false);
expect(result.totalBytes).toBe(byteLength(content));
expect(result.outputBytes).toBe(byteLength(content));
expect(result.totalBytes).toBe(9);
});
it("truncates head on UTF-8 byte limits without partial lines", () => {
const content = "éé\nabc";
const result = truncateHead(content, { maxBytes: 4, maxLines: 10 });
expect(result.content).toBe("éé");
expect(result.truncated).toBe(true);
expect(result.truncatedBy).toBe("bytes");
expect(result.outputBytes).toBe(4);
expect(result.firstLineExceedsLimit).toBe(false);
});
it("reports head truncation when the first line exceeds the byte limit", () => {
const result = truncateHead("éé\nabc", { maxBytes: 3, maxLines: 10 });
expect(result.content).toBe("");
expect(result.truncated).toBe(true);
expect(result.truncatedBy).toBe("bytes");
expect(result.firstLineExceedsLimit).toBe(true);
});
it("truncates tail on UTF-8 boundaries when only a partial last line fits", () => {
const result = truncateTail("aé🙂b", { maxBytes: 5, maxLines: 10 });
expect(result.content).toBe("🙂b");
expect(result.truncated).toBe(true);
expect(result.truncatedBy).toBe("bytes");
expect(result.lastLinePartial).toBe(true);
expect(result.outputBytes).toBe(5);
});
it("drops an oversized trailing character when it cannot fit in tail byte limit", () => {
const result = truncateTail("abc🙂", { maxBytes: 3, maxLines: 10 });
expect(result.content).toBe("");
expect(result.truncated).toBe(true);
expect(result.truncatedBy).toBe("bytes");
expect(result.lastLinePartial).toBe(true);
expect(result.outputBytes).toBe(0);
});
it("matches Buffer tail truncation semantics for surrogate edge cases", () => {
const inputs = ["a\ud83d", "\ude42b", "a\ude42b", "\ud83d\ud83d\ude42", "\ud83d\ude42\ude42", "👩‍💻"];
for (const input of inputs) assertMatchesBufferTail(input);
});
it("matches Buffer tail truncation semantics across deterministic fuzz cases", () => {
const alphabet = [
"a",
"\u007f",
"\u0080",
"é",
"\u07ff",
"\u0800",
"中",
"\ud7ff",
"\ud800",
"\ud83d",
"\udc00",
"\ude42",
"🙂",
"\ue000",
"\uffff",
];
function checkExhaustive(prefix: string, depth: number): void {
assertMatchesBufferTail(prefix, sampledByteLimits(prefix));
if (depth === 0) return;
for (const character of alphabet) checkExhaustive(prefix + character, depth - 1);
}
checkExhaustive("", 3);
let seed = 0x12345678;
function random(): number {
seed = (seed * 1664525 + 1013904223) >>> 0;
return seed / 0x100000000;
}
for (let i = 0; i < 1_000; i++) {
let input = "";
const length = Math.floor(random() * 80);
for (let j = 0; j < length; j++) input += alphabet[Math.floor(random() * alphabet.length)];
assertMatchesBufferTail(input, sampledByteLimits(input));
}
});
});