feat(agent): add harness stream configuration

This commit is contained in:
Mario Zechner
2026-05-10 18:12:59 +02:00
parent f8d0fa67dc
commit c0f416aabc
3 changed files with 371 additions and 35 deletions

View File

@@ -1,4 +1,10 @@
import type { AssistantMessage, ImageContent, Model, UserMessage } from "@earendil-works/pi-ai";
import {
type AssistantMessage,
type ImageContent,
type Model,
streamSimple,
type UserMessage,
} from "@earendil-works/pi-ai";
import { Agent, type QueueMode } from "../agent.js";
import type { AgentEvent, AgentMessage, AgentTool, ThinkingLevel } from "../types.js";
import { collectEntriesForBranchSummary, generateBranchSummary } from "./compaction/branch-summarization.js";
@@ -13,7 +19,8 @@ import type {
AgentHarnessOwnEvent,
AgentHarnessPhase,
AgentHarnessResources,
AgentHarnessTurnState,
AgentHarnessStreamOptions,
AgentHarnessStreamOptionsPatch,
ExecutionEnv,
NavigateTreeResult,
PendingSessionWrite,
@@ -28,6 +35,87 @@ function createUserMessage(text: string, images?: ImageContent[]): UserMessage {
return { role: "user", content, timestamp: Date.now() };
}
function cloneStreamOptions(streamOptions?: AgentHarnessStreamOptions): AgentHarnessStreamOptions {
return {
...streamOptions,
headers: streamOptions?.headers ? { ...streamOptions.headers } : undefined,
metadata: streamOptions?.metadata ? { ...streamOptions.metadata } : undefined,
};
}
function mergeHeaders(...headers: Array<Record<string, string> | undefined>): Record<string, string> | undefined {
const merged: Record<string, string> = {};
let hasHeaders = false;
for (const entry of headers) {
if (!entry) continue;
Object.assign(merged, entry);
hasHeaders = true;
}
return hasHeaders ? merged : undefined;
}
function hasOwn(object: object, key: PropertyKey): boolean {
return Object.hasOwn(object, key);
}
function applyStreamOptionsPatch(
base: AgentHarnessStreamOptions,
patch?: AgentHarnessStreamOptionsPatch,
): AgentHarnessStreamOptions {
const result = cloneStreamOptions(base);
if (!patch) return result;
if (hasOwn(patch, "transport")) result.transport = patch.transport;
if (hasOwn(patch, "timeoutMs")) result.timeoutMs = patch.timeoutMs;
if (hasOwn(patch, "maxRetries")) result.maxRetries = patch.maxRetries;
if (hasOwn(patch, "maxRetryDelayMs")) result.maxRetryDelayMs = patch.maxRetryDelayMs;
if (hasOwn(patch, "cacheRetention")) result.cacheRetention = patch.cacheRetention;
if (hasOwn(patch, "headers")) {
if (patch.headers === undefined) {
result.headers = undefined;
} else {
const headers = { ...(result.headers ?? {}) };
for (const [key, value] of Object.entries(patch.headers)) {
if (value === undefined) delete headers[key];
else headers[key] = value;
}
result.headers = Object.keys(headers).length > 0 ? headers : undefined;
}
}
if (hasOwn(patch, "metadata")) {
if (patch.metadata === undefined) {
result.metadata = undefined;
} else {
const metadata = { ...(result.metadata ?? {}) };
for (const [key, value] of Object.entries(patch.metadata)) {
if (value === undefined) delete metadata[key];
else metadata[key] = value;
}
result.metadata = Object.keys(metadata).length > 0 ? metadata : undefined;
}
}
return result;
}
interface AgentHarnessTurnState<
TSkill extends Skill = Skill,
TPromptTemplate extends PromptTemplate = PromptTemplate,
TTool extends AgentTool = AgentTool,
> {
messages: AgentMessage[];
resources: AgentHarnessResources<TSkill, TPromptTemplate>;
streamOptions: AgentHarnessStreamOptions;
sessionId: string;
systemPrompt: string;
model: Model<any>;
thinkingLevel: ThinkingLevel;
tools: TTool[];
activeTools: TTool[];
}
export class AgentHarness<
TSkill extends Skill = Skill,
TPromptTemplate extends PromptTemplate = PromptTemplate,
@@ -45,6 +133,9 @@ export class AgentHarness<
private followUpQueue: UserMessage[] = [];
private pendingSessionWrites: PendingSessionWrite[] = [];
private resources: AgentHarnessResources<TSkill, TPromptTemplate>;
private streamOptions: AgentHarnessStreamOptions;
private appliedStreamOptions: AgentHarnessStreamOptions = {};
private appliedSessionId?: string;
private systemPrompt: AgentHarnessOptions<TSkill, TPromptTemplate, TTool>["systemPrompt"];
private getApiKeyAndHeaders?: AgentHarnessOptions["getApiKeyAndHeaders"];
private tools = new Map<string, TTool>();
@@ -60,12 +151,46 @@ export class AgentHarness<
thinkingLevel: options.thinkingLevel,
tools: options.tools ?? [],
},
streamFn: async (model, context, streamOptions) => {
const auth = await this.getApiKeyAndHeaders?.(model);
const snapshotOptions: AgentHarnessStreamOptions = {
...this.appliedStreamOptions,
headers: mergeHeaders(this.appliedStreamOptions.headers, auth?.headers),
};
const requestOptions = await this.emitBeforeProviderRequest(
model,
this.appliedSessionId ?? "",
snapshotOptions,
);
return streamSimple(model, context, {
cacheRetention: requestOptions.cacheRetention,
headers: requestOptions.headers,
maxRetries: requestOptions.maxRetries,
maxRetryDelayMs: requestOptions.maxRetryDelayMs,
metadata: requestOptions.metadata,
onPayload: async (payload) => await this.emitBeforeProviderPayload(model, payload),
onResponse: async (response) => {
const headers = { ...(response.headers as Record<string, string>) };
await this.emitOwn(
{ type: "after_provider_response", status: response.status, headers },
this.agent.signal,
);
},
reasoning: streamOptions?.reasoning,
signal: streamOptions?.signal,
sessionId: this.appliedSessionId,
timeoutMs: requestOptions.timeoutMs,
transport: requestOptions.transport,
apiKey: auth?.apiKey,
});
},
steeringMode: options.steeringMode,
followUpMode: options.followUpMode,
});
this.env = options.env;
this.session = options.session;
this.resources = options.resources ?? {};
this.streamOptions = cloneStreamOptions(options.streamOptions);
this.systemPrompt = options.systemPrompt;
this.getApiKeyAndHeaders = options.getApiKeyAndHeaders;
for (const tool of options.tools ?? []) {
@@ -76,11 +201,6 @@ export class AgentHarness<
this.activeToolNames = options.activeToolNames ?? (options.tools ?? []).map((tool) => tool.name);
this.agent.state.model = this.model;
this.agent.state.thinkingLevel = this.thinkingLevel;
this.agent.getApiKey = async (provider) => {
const model = this.model;
if (!this.getApiKeyAndHeaders || model.provider !== provider) return undefined;
return (await this.getApiKeyAndHeaders(model))?.apiKey;
};
this.agent.transformContext = async (messages) => {
const result = await this.emitHook({ type: "context", messages: [...messages] });
return result?.messages ?? messages;
@@ -108,14 +228,6 @@ export class AgentHarness<
? { content: patch.content, details: patch.details, isError: patch.isError, terminate: patch.terminate }
: undefined;
};
this.agent.onPayload = async (payload) => {
const result = await this.emitHook({ type: "before_provider_request", payload });
return result?.payload ?? payload;
};
this.agent.onResponse = async (response) => {
const headers = { ...(response.headers as Record<string, string>) };
await this.emitOwn({ type: "after_provider_response", status: response.status, headers }, this.agent.signal);
};
this.agent.prepareNextTurn = async () => {
await this.flushPendingSessionWrites();
const turnState = await this.createTurnState();
@@ -162,6 +274,41 @@ export class AgentHarness<
return lastResult;
}
private async emitBeforeProviderRequest(
model: Model<any>,
sessionId: string,
streamOptions: AgentHarnessStreamOptions,
): Promise<AgentHarnessStreamOptions> {
const handlers = this.hooks.get("before_provider_request");
let current = cloneStreamOptions(streamOptions);
if (!handlers || handlers.size === 0) return current;
for (const handler of handlers) {
const result = await handler({
type: "before_provider_request",
model,
sessionId,
streamOptions: cloneStreamOptions(current),
});
if (result?.streamOptions) {
current = applyStreamOptionsPatch(current, result.streamOptions);
}
}
return current;
}
private async emitBeforeProviderPayload(model: Model<any>, payload: unknown): Promise<unknown> {
const handlers = this.hooks.get("before_provider_payload");
let current = payload;
if (!handlers || handlers.size === 0) return current;
for (const handler of handlers) {
const result = await handler({ type: "before_provider_payload", model, payload: current });
if (result !== undefined) {
current = result.payload;
}
}
return current;
}
private async emitQueueUpdate(): Promise<void> {
await this.emitOwn({
type: "queue_update",
@@ -174,6 +321,7 @@ export class AgentHarness<
private async createTurnState(): Promise<AgentHarnessTurnState<TSkill, TPromptTemplate, TTool>> {
const context = await this.session.buildContext();
const resources = this.getResources();
const sessionMetadata = await this.session.getMetadata();
const tools = [...this.tools.values()];
const activeTools = this.activeToolNames
.map((name) => this.tools.get(name))
@@ -194,6 +342,8 @@ export class AgentHarness<
return {
messages: context.messages,
resources,
streamOptions: cloneStreamOptions(this.streamOptions),
sessionId: sessionMetadata.id,
systemPrompt,
model: this.model,
thinkingLevel: this.thinkingLevel,
@@ -204,6 +354,8 @@ export class AgentHarness<
private applyTurnState(turnState: AgentHarnessTurnState<TSkill, TPromptTemplate, TTool>): void {
this.agent.state.messages = turnState.messages;
this.appliedStreamOptions = cloneStreamOptions(turnState.streamOptions);
this.appliedSessionId = turnState.sessionId;
this.agent.state.systemPrompt = turnState.systemPrompt;
this.agent.state.model = turnState.model;
this.agent.state.thinkingLevel = turnState.thinkingLevel;
@@ -602,6 +754,14 @@ export class AgentHarness<
await this.emitOwn({ type: "resources_update", resources: this.getResources(), previousResources });
}
getStreamOptions(): AgentHarnessStreamOptions {
return cloneStreamOptions(this.streamOptions);
}
setStreamOptions(streamOptions: AgentHarnessStreamOptions): void {
this.streamOptions = cloneStreamOptions(streamOptions);
}
async setTools(tools: TTool[], activeToolNames?: string[]): Promise<void> {
this.tools = new Map(tools.map((tool) => [tool.name, tool]));
if (activeToolNames) {

View File

@@ -1,4 +1,4 @@
import type { ImageContent, Model, TextContent } from "@earendil-works/pi-ai";
import type { ImageContent, Model, SimpleStreamOptions, TextContent, Transport } from "@earendil-works/pi-ai";
import type { QueueMode } from "../agent.js";
import type { AgentEvent, AgentMessage, AgentTool, ThinkingLevel } from "../index.js";
import type { Session } from "./session/session.js";
@@ -43,6 +43,33 @@ export interface AgentHarnessResources<
skills?: TSkill[];
}
/** Curated provider request options owned by the harness and snapshotted per turn. */
export interface AgentHarnessStreamOptions {
/** Preferred transport forwarded to the stream function. */
transport?: Transport;
/** Provider request timeout in milliseconds. */
timeoutMs?: number;
/** Maximum provider retry attempts. */
maxRetries?: number;
/** Optional cap for provider-requested retry delays. */
maxRetryDelayMs?: number;
/** Additional request headers merged with auth and lifecycle headers. */
headers?: Record<string, string>;
/** Provider metadata forwarded with requests. */
metadata?: SimpleStreamOptions["metadata"];
/** Provider cache retention hint. */
cacheRetention?: SimpleStreamOptions["cacheRetention"];
}
/** Per-request stream option patch returned by provider hooks. */
export interface AgentHarnessStreamOptionsPatch
extends Omit<Partial<AgentHarnessStreamOptions>, "headers" | "metadata"> {
/** Header patch. `undefined` values delete keys; explicit `headers: undefined` clears all headers. */
headers?: Record<string, string | undefined>;
/** Metadata patch. `undefined` values delete keys; explicit `metadata: undefined` clears all metadata. */
metadata?: Record<string, unknown | undefined>;
}
/** Kind of filesystem object as addressed by an {@link ExecutionEnv}. Symlinks are not followed automatically. */
export type FileKind = "file" | "directory" | "symlink";
@@ -298,20 +325,6 @@ export type PendingSessionWrite = SessionTreeEntry extends infer TEntry
: never
: never;
export interface AgentHarnessTurnState<
TSkill extends Skill = Skill,
TPromptTemplate extends PromptTemplate = PromptTemplate,
TTool extends AgentTool = AgentTool,
> {
messages: AgentMessage[];
resources: AgentHarnessResources<TSkill, TPromptTemplate>;
systemPrompt: string;
model: Model<any>;
thinkingLevel: ThinkingLevel;
tools: TTool[];
activeTools: TTool[];
}
export interface QueueUpdateEvent {
type: "queue_update";
steer: AgentMessage[];
@@ -353,6 +366,14 @@ export interface ContextEvent {
export interface BeforeProviderRequestEvent {
type: "before_provider_request";
model: Model<any>;
sessionId: string;
streamOptions: AgentHarnessStreamOptions;
}
export interface BeforeProviderPayloadEvent {
type: "before_provider_payload";
model: Model<any>;
payload: unknown;
}
@@ -440,6 +461,7 @@ export type AgentHarnessOwnEvent<
| BeforeAgentStartEvent<TSkill, TPromptTemplate>
| ContextEvent
| BeforeProviderRequestEvent
| BeforeProviderPayloadEvent
| AfterProviderResponseEvent
| ToolCallEvent
| ToolResultEvent
@@ -465,6 +487,10 @@ export interface ContextResult {
}
export interface BeforeProviderRequestResult {
streamOptions?: AgentHarnessStreamOptionsPatch;
}
export interface BeforeProviderPayloadResult {
payload: unknown;
}
@@ -497,6 +523,7 @@ export type AgentHarnessEventResultMap = {
before_agent_start: BeforeAgentStartResult | undefined;
context: ContextResult | undefined;
before_provider_request: BeforeProviderRequestResult | undefined;
before_provider_payload: BeforeProviderPayloadResult | undefined;
after_provider_response: undefined;
tool_call: ToolCallResult | undefined;
tool_result: ToolResultPatch | undefined;
@@ -613,6 +640,8 @@ export interface AgentHarnessOptions<
getApiKeyAndHeaders?: (
model: Model<any>,
) => Promise<{ apiKey: string; headers?: Record<string, string> } | undefined>;
/** Curated stream/provider request options. Snapshotted at turn start. */
streamOptions?: AgentHarnessStreamOptions;
model: Model<any>;
thinkingLevel?: ThinkingLevel;
activeToolNames?: string[];