Skip to content
Open
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
25 changes: 24 additions & 1 deletion packages/sandbox-image/src/pi-driver.ts
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@ import type {
ConsolidationResult,
PlanResult,
TurnResult,
ThinkingLevel,
} from "./runtime.js";

const DEFAULT_CODEVIL_PI_AGENT_DIR = "/opt/codevil/pi-agent";
Expand Down Expand Up @@ -137,6 +138,7 @@ export class PiAgentDriver implements AgentDriver {
cwd: options.cwd,
agentDir,
model: proxiedModel,
...(options.thinkingLevel ? { thinkingLevel: options.thinkingLevel } : {}),
authStorage,
modelRegistry,
customTools,
Expand Down Expand Up @@ -282,8 +284,14 @@ export class PiAgentDriver implements AgentDriver {
}
}

async switchToExecution(modelId: string, provider = "anthropic"): Promise<void> {
async switchToExecution(
modelId: string,
provider = "anthropic",
thinkingLevel?: ThinkingLevel,
): Promise<void> {
const session = this.requireSession();
this.assertSessionIdle("change model or reasoning effort");
const nextThinkingLevel = thinkingLevel ?? session.thinkingLevel;
const modelRegistry = this.modelRegistry;
if (!modelRegistry) throw new Error("Model registry has not been initialized");
const authStorage = this.authStorage;
Expand All @@ -307,6 +315,15 @@ export class PiAgentDriver implements AgentDriver {

session.setActiveToolsByName(["read", "bash", "edit", "write"]);
await session.setModel(proxiedModel);
// Apply after setModel so Pi can clamp the requested level to the new model.
// This preserves the previous effort when only the model changes.
session.setThinkingLevel(nextThinkingLevel);
}

setThinkingLevel(thinkingLevel: ThinkingLevel): void {
const session = this.requireSession();
this.assertSessionIdle("change reasoning effort");
session.setThinkingLevel(thinkingLevel);
}

refreshProxyCapabilities(tokens: Partial<Record<ProviderApi, string>>): void {
Expand All @@ -332,6 +349,12 @@ export class PiAgentDriver implements AgentDriver {
if (!this.session) throw new Error("Pi session has not been started");
return this.session;
}

private assertSessionIdle(operation: string): void {
if (this.session?.isStreaming) {
throw new Error(`Cannot ${operation} while an agent turn is streaming; retry between turns.`);
}
}
}

async function createCodevilResourceLoader(
Expand Down
14 changes: 13 additions & 1 deletion packages/sandbox-image/src/runtime-types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -13,11 +13,16 @@ import type { CommandRunner, Verifier } from "./verification.js";

export type ProxyCapabilities = Partial<Record<ProviderApi, string>> & { git?: string };

/** Pi's provider-neutral reasoning/effort levels. */
export type ThinkingLevel = "off" | "minimal" | "low" | "medium" | "high" | "xhigh";

export interface AgentStartOptions {
cwd: string;
mode: "coding";
model: string;
provider: string;
/** Reasoning effort for the first turn; omitted to use Pi's default. */
thinkingLevel?: ThinkingLevel;
providerConfig?: ProviderPublicConfig;
llmKey?: string;
proxyBase?: string;
Expand Down Expand Up @@ -90,7 +95,14 @@ export interface AgentDriver {
consolidateAnnotations?(input: ConsolidationInput): Promise<ConsolidationResult>;
/** Replace expiring sandbox-only proxy capabilities without exposing provider keys. */
refreshProxyCapabilities?(tokens: Partial<Record<ProviderApi, string>>): Promise<void> | void;
switchToExecution(model: string, provider?: string): Promise<void>;
/**
* Change model and, optionally, reasoning effort on the existing conversation.
* Call only between turns so no provider request is interrupted.
*/
switchToExecution(model: string, provider?: string, thinkingLevel?: ThinkingLevel): Promise<void>;
/** Change reasoning effort on the existing conversation between turns. */
setThinkingLevel?(thinkingLevel: ThinkingLevel): void;

execute(plan: string): Promise<CostInfo>;
dispose?(): Promise<void> | void;
}
Expand Down
25 changes: 19 additions & 6 deletions packages/sandbox-image/src/runtime.ts
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@ export type {
ConsolidationResult,
AgentDriver,
AgentDriverFactory,
ThinkingLevel,
GitDriver,
ProxyCapabilities,
PushBranchOptions,
Expand All @@ -69,6 +70,7 @@ import type {
AgentDriverFactory,
AskQuestionOutcome,
AskQuestionParams,
ThinkingLevel,
CreatePullRequestToolOptions,
GitDriver,
ProxyCapabilities,
Expand Down Expand Up @@ -161,19 +163,19 @@ export class SandboxRuntime {
await this.handleInit(message.repo, message.restored_from_cache ?? false);
return;
case "agent_turn":
await this.handleAgentTurn(message.run_id, message.prompt, message.model, message.provider, parent);
await this.handleAgentTurn(message.run_id, message.prompt, message.model, message.provider, message.thinking_level, parent);
return;
case "plan":
await this.handlePlan(message.run_id, message.prompt, message.model, message.provider, parent);
await this.handlePlan(message.run_id, message.prompt, message.model, message.provider, message.thinking_level, parent);
return;
case "refine_plan":
await this.handleRefine(message.feedback, parent);
await this.handleRefine(message.feedback, message.thinking_level, parent);
return;
case "consolidate_annotations":
await this.handleConsolidateAnnotations(message, parent);
return;
case "execute":
await this.handleExecute(message.plan, message.model, message.provider, parent);
await this.handleExecute(message.plan, message.model, message.provider, message.thinking_level, parent);
return;
case "create_pr":
await this.handleCreatePullRequest(message, parent);
Expand Down Expand Up @@ -292,6 +294,7 @@ export class SandboxRuntime {
prompt: string,
model: string,
provider: string | undefined,
thinkingLevel: ThinkingLevel | undefined,
parent: SpanContext | undefined,
): Promise<void> {
const repoDir = this.requireRepo().dir;
Expand All @@ -303,6 +306,7 @@ export class SandboxRuntime {
mode: "coding",
model,
provider: provider ?? this.provider,
...(thinkingLevel ? { thinkingLevel } : {}),
providerConfig: this.providerConfig,
llmKey: this.llmKey,
proxyBase: this.proxyBase,
Expand Down Expand Up @@ -339,6 +343,7 @@ export class SandboxRuntime {
prompt: string,
model: string,
provider: string | undefined,
thinkingLevel: ThinkingLevel | undefined,
parent: SpanContext | undefined,
): Promise<void> {
const repoDir = this.requireRepo().dir;
Expand All @@ -349,6 +354,7 @@ export class SandboxRuntime {
mode: "coding",
model,
provider: provider ?? this.provider,
...(thinkingLevel ? { thinkingLevel } : {}),
providerConfig: this.providerConfig,
llmKey: this.llmKey,
proxyBase: this.proxyBase,
Expand All @@ -364,6 +370,7 @@ export class SandboxRuntime {
this.agent = agent;
}

if (thinkingLevel !== undefined) this.agent.setThinkingLevel?.(thinkingLevel);
this.activeRunId = runId;
try {
try {
Expand Down Expand Up @@ -421,8 +428,13 @@ export class SandboxRuntime {
});
}

private async handleRefine(feedback: string, parent: SpanContext | undefined): Promise<void> {
private async handleRefine(
feedback: string,
thinkingLevel: ThinkingLevel | undefined,
parent: SpanContext | undefined,
): Promise<void> {
const agent = this.requireAgent();
if (thinkingLevel !== undefined) agent.setThinkingLevel?.(thinkingLevel);
const result = await this.maybeSpan("llm.refine", { parent }, () =>
agent.refine(refinePrompt(feedback)),
);
Expand Down Expand Up @@ -486,10 +498,11 @@ export class SandboxRuntime {
plan: string,
model: string,
provider: string | undefined,
thinkingLevel: ThinkingLevel | undefined,
parent: SpanContext | undefined,
): Promise<void> {
const agent = this.requireAgent();
await agent.switchToExecution(model, provider ?? this.provider);
await agent.switchToExecution(model, provider ?? this.provider, thinkingLevel);
let cost = await this.maybeSpan(
"llm.execute",
{ parent, attributes: { model, provider: provider ?? this.provider } },
Expand Down
41 changes: 41 additions & 0 deletions packages/sandbox-image/test/pi-driver.test.mjs
Original file line number Diff line number Diff line change
Expand Up @@ -272,6 +272,47 @@ test("Cloudflare AI Gateway uses a proxy capability and its resolved account/gat
}
});

test("changes reasoning effort and model on the existing conversation without rewriting messages", async () => {
const cwd = await mkdtemp(join(tmpdir(), "codevil-pi-effort-switch-"));
const driver = new PiAgentDriver();
try {
await driver.start({
cwd,
mode: "coding",
provider: "anthropic",
model: "claude-haiku-4-5",
thinkingLevel: "low",
proxyBase: "https://worker.example",
proxySessionId: "ses_effort_switch",
proxyTokens: { "anthropic-messages": "initial-capability" },
onEvent: () => {},
createPullRequest: async () => ({ url: "https://github.com/example/app/pull/1" }),
});

const messagesBefore = driver.session.messages;
assert.equal(driver.session.thinkingLevel, "low");
driver.setThinkingLevel("high");
await driver.switchToExecution("claude-haiku-4-5", "anthropic", "minimal");

assert.equal(driver.session.thinkingLevel, "minimal");
assert.equal(driver.session.messages, messagesBefore);
assert.equal(driver.session.model.id, "claude-haiku-4-5");
} finally {
driver.dispose();
await rm(cwd, { recursive: true, force: true });
}
});

test("rejects model or effort changes while a provider request is streaming", () => {
const driver = new PiAgentDriver();
driver.session = { isStreaming: true };

assert.throws(
() => driver.setThinkingLevel("high"),
/while an agent turn is streaming/,
);
});

test("switchToExecution preserves the Pi provider target and refreshes the selected API capability", async () => {
const cwd = await mkdtemp(join(tmpdir(), "codevil-pi-proxy-switch-"));
const driver = new PiAgentDriver();
Expand Down
1 change: 1 addition & 0 deletions packages/shared/src/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -251,6 +251,7 @@ export type {
SandboxToDOMessage,
} from "./messages-sandbox.js";
export {
ThinkingLevelSchema,
InitMessageSchema,
AgentTurnMessageSchema,
PlanMessageSchema,
Expand Down
7 changes: 7 additions & 0 deletions packages/shared/src/messages-sandbox.ts
Original file line number Diff line number Diff line change
Expand Up @@ -35,12 +35,16 @@ const TraceContextFields = {
parent_span_id: z.string().optional(),
};

/** Provider-neutral reasoning effort understood by the sandbox agent. */
export const ThinkingLevelSchema = z.enum(["off", "minimal", "low", "medium", "high", "xhigh"]);

export const PlanMessageSchema = z.object({
type: z.literal("plan"),
run_id: z.string(),
prompt: z.string(),
model: z.string(),
provider: z.string().optional(),
thinking_level: ThinkingLevelSchema.optional(),
...TraceContextFields,
});

Expand All @@ -50,6 +54,7 @@ export const AgentTurnMessageSchema = z.object({
prompt: z.string(),
model: z.string(),
provider: z.string().optional(),
thinking_level: ThinkingLevelSchema.optional(),
...TraceContextFields,
});

Expand All @@ -58,12 +63,14 @@ export const ExecuteMessageSchema = z.object({
plan: z.string(),
model: z.string(),
provider: z.string().optional(),
thinking_level: ThinkingLevelSchema.optional(),
...TraceContextFields,
});

export const RefinePlanSandboxMessageSchema = z.object({
type: z.literal("refine_plan"),
feedback: z.string(),
thinking_level: ThinkingLevelSchema.optional(),
...TraceContextFields,
});

Expand Down
Loading