feat(DIS-153): model selection per agent #62
7 changed files with 254 additions and 13 deletions
|
|
@ -4,6 +4,8 @@ import * as YAML from "yaml";
|
||||||
import { ZodError } from "zod";
|
import { ZodError } from "zod";
|
||||||
import { AgentYamlSchema } from "./schema";
|
import { AgentYamlSchema } from "./schema";
|
||||||
import type { AgentYaml as AgentIdentity } from "./schema";
|
import type { AgentYaml as AgentIdentity } from "./schema";
|
||||||
|
import { DEFAULT_MODEL } from "../config/models";
|
||||||
|
import type { AllowedModel } from "../config/models";
|
||||||
|
|
||||||
export type { AgentYaml as AgentIdentity } from "./schema";
|
export type { AgentYaml as AgentIdentity } from "./schema";
|
||||||
|
|
||||||
|
|
@ -44,7 +46,8 @@ export function createAgentYaml(
|
||||||
workspacePath: string,
|
workspacePath: string,
|
||||||
name: string,
|
name: string,
|
||||||
role: string,
|
role: string,
|
||||||
channelId: string
|
channelId: string,
|
||||||
|
model: AllowedModel = DEFAULT_MODEL
|
||||||
): void {
|
): void {
|
||||||
const identity: AgentIdentity = {
|
const identity: AgentIdentity = {
|
||||||
name,
|
name,
|
||||||
|
|
@ -54,6 +57,7 @@ export function createAgentYaml(
|
||||||
.join(" "),
|
.join(" "),
|
||||||
role,
|
role,
|
||||||
channel_id: channelId,
|
channel_id: channelId,
|
||||||
|
model,
|
||||||
};
|
};
|
||||||
|
|
||||||
const yamlContent = YAML.stringify(identity);
|
const yamlContent = YAML.stringify(identity);
|
||||||
|
|
@ -169,13 +173,14 @@ export function setupAgentWorkspace(
|
||||||
workspacePath: string,
|
workspacePath: string,
|
||||||
name: string,
|
name: string,
|
||||||
role: string,
|
role: string,
|
||||||
channelId: string
|
channelId: string,
|
||||||
|
model: AllowedModel = DEFAULT_MODEL
|
||||||
): void {
|
): void {
|
||||||
// Ensure workspace directory exists
|
// Ensure workspace directory exists
|
||||||
fs.mkdirSync(workspacePath, { recursive: true });
|
fs.mkdirSync(workspacePath, { recursive: true });
|
||||||
|
|
||||||
// Create all required files
|
// Create all required files
|
||||||
createAgentYaml(workspacePath, name, role, channelId);
|
createAgentYaml(workspacePath, name, role, channelId, model);
|
||||||
createClaudeMd(workspacePath, name, role);
|
createClaudeMd(workspacePath, name, role);
|
||||||
createClaudeConfig(workspacePath);
|
createClaudeConfig(workspacePath);
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@ import { resolveClaude } from "../runtime/resolve-claude";
|
||||||
import { sanitizedEnv } from "../runtime/env";
|
import { sanitizedEnv } from "../runtime/env";
|
||||||
import { Semaphore } from "../runtime/concurrency";
|
import { Semaphore } from "../runtime/concurrency";
|
||||||
import type { RunResult } from "./types";
|
import type { RunResult } from "./types";
|
||||||
|
import { DEFAULT_MODEL } from "../config/models";
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
// Module-level semaphore — lazy-initialised on first runAgent() call so the
|
// Module-level semaphore — lazy-initialised on first runAgent() call so the
|
||||||
|
|
@ -206,11 +207,14 @@ export async function runAgent(options: RunAgentOptions): Promise<RunResult> {
|
||||||
// ~/.disclaw/workspaces/ (outside the DisClaw repo), so walk-up will never
|
// ~/.disclaw/workspaces/ (outside the DisClaw repo), so walk-up will never
|
||||||
// reach the DisClaw project CLAUDE.md. The workspace CLAUDE.md is picked up
|
// reach the DisClaw project CLAUDE.md. The workspace CLAUDE.md is picked up
|
||||||
// automatically because cwd is set to workspacePath.
|
// automatically because cwd is set to workspacePath.
|
||||||
|
const model = identity.model ?? DEFAULT_MODEL;
|
||||||
const args = [
|
const args = [
|
||||||
"-p",
|
"-p",
|
||||||
prompt,
|
prompt,
|
||||||
"--output-format",
|
"--output-format",
|
||||||
"json",
|
"json",
|
||||||
|
"--model",
|
||||||
|
model,
|
||||||
];
|
];
|
||||||
|
|
||||||
const child = spawn(claudeCommand, args, {
|
const child = spawn(claudeCommand, args, {
|
||||||
|
|
|
||||||
|
|
@ -1,10 +1,14 @@
|
||||||
import { z } from "zod";
|
import { z } from "zod";
|
||||||
|
import { ALLOWED_MODELS, DEFAULT_MODEL } from "../config/models";
|
||||||
|
|
||||||
export const AgentYamlSchema = z.object({
|
export const AgentYamlSchema = z.object({
|
||||||
name: z.string().min(1),
|
name: z.string().min(1),
|
||||||
display_name: z.string().optional(),
|
display_name: z.string().optional(),
|
||||||
role: z.string().optional(),
|
role: z.string().optional(),
|
||||||
channel_id: z.string().min(1),
|
channel_id: z.string().min(1),
|
||||||
|
model: z
|
||||||
|
.enum(ALLOWED_MODELS.map((m) => m.id) as [string, ...string[]])
|
||||||
|
.default(DEFAULT_MODEL),
|
||||||
});
|
});
|
||||||
|
|
||||||
export type AgentYaml = z.infer<typeof AgentYamlSchema>;
|
export type AgentYaml = z.infer<typeof AgentYamlSchema>;
|
||||||
|
|
|
||||||
|
|
@ -3,11 +3,17 @@ import {
|
||||||
ChannelType,
|
ChannelType,
|
||||||
SlashCommandBuilder,
|
SlashCommandBuilder,
|
||||||
Guild,
|
Guild,
|
||||||
|
ActionRowBuilder,
|
||||||
|
StringSelectMenuBuilder,
|
||||||
|
StringSelectMenuInteraction,
|
||||||
|
ComponentType,
|
||||||
} from "discord.js";
|
} from "discord.js";
|
||||||
import * as path from "path";
|
import * as path from "path";
|
||||||
import { DisclawDatabase } from "../db/database";
|
import { DisclawDatabase } from "../db/database";
|
||||||
import { DisclawConfig } from "../config/loader";
|
import { DisclawConfig } from "../config/loader";
|
||||||
import { setupAgentWorkspace } from "../agent/identity";
|
import { setupAgentWorkspace } from "../agent/identity";
|
||||||
|
import { ALLOWED_MODELS, DEFAULT_MODEL } from "../config/models";
|
||||||
|
import type { AllowedModel } from "../config/models";
|
||||||
|
|
||||||
// Allowed characters for agent names: lowercase letters, numbers, hyphens
|
// Allowed characters for agent names: lowercase letters, numbers, hyphens
|
||||||
const NAME_PATTERN = /^[a-z0-9][a-z0-9-]{0,30}[a-z0-9]$/;
|
const NAME_PATTERN = /^[a-z0-9][a-z0-9-]{0,30}[a-z0-9]$/;
|
||||||
|
|
@ -82,7 +88,50 @@ export async function handleNewAgent(
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
await interaction.deferReply();
|
// Model selection via ephemeral select menu (before deferring the final reply)
|
||||||
|
const modelSelect = new StringSelectMenuBuilder()
|
||||||
|
.setCustomId("model-select")
|
||||||
|
.setPlaceholder("Modell auswählen …")
|
||||||
|
.addOptions(
|
||||||
|
ALLOWED_MODELS.map((m) => ({
|
||||||
|
label: m.label,
|
||||||
|
value: m.id,
|
||||||
|
default: m.id === DEFAULT_MODEL,
|
||||||
|
}))
|
||||||
|
);
|
||||||
|
|
||||||
|
const row = new ActionRowBuilder<StringSelectMenuBuilder>().addComponents(
|
||||||
|
modelSelect
|
||||||
|
);
|
||||||
|
|
||||||
|
await interaction.reply({
|
||||||
|
content: `Welches Modell soll **${name}** verwenden?`,
|
||||||
|
components: [row],
|
||||||
|
ephemeral: true,
|
||||||
|
});
|
||||||
|
|
||||||
|
let selectedModel: AllowedModel = DEFAULT_MODEL;
|
||||||
|
|
||||||
|
try {
|
||||||
|
const selectInteraction = await interaction.channel!.awaitMessageComponent<ComponentType.StringSelect>({
|
||||||
|
componentType: ComponentType.StringSelect,
|
||||||
|
filter: (i: StringSelectMenuInteraction) =>
|
||||||
|
i.customId === "model-select" && i.user.id === interaction.user.id,
|
||||||
|
time: 60_000,
|
||||||
|
});
|
||||||
|
|
||||||
|
selectedModel = selectInteraction.values[0] as AllowedModel;
|
||||||
|
await selectInteraction.update({
|
||||||
|
content: `Modell **${selectedModel}** ausgewählt. Erstelle Agenten …`,
|
||||||
|
components: [],
|
||||||
|
});
|
||||||
|
} catch {
|
||||||
|
// Timeout or no selection — use default, update the reply
|
||||||
|
await interaction.editReply({
|
||||||
|
content: `Keine Auswahl — verwende Standardmodell **${DEFAULT_MODEL}**. Erstelle Agenten …`,
|
||||||
|
components: [],
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
try {
|
try {
|
||||||
// 1. Create Discord channel
|
// 1. Create Discord channel
|
||||||
|
|
@ -95,17 +144,20 @@ export async function handleNewAgent(
|
||||||
|
|
||||||
// 2. Create workspace with full Claude Code environment
|
// 2. Create workspace with full Claude Code environment
|
||||||
const workspacePath = path.resolve(config.workspaces_root, name);
|
const workspacePath = path.resolve(config.workspaces_root, name);
|
||||||
setupAgentWorkspace(workspacePath, name, role, channel.id);
|
setupAgentWorkspace(workspacePath, name, role, channel.id, selectedModel);
|
||||||
|
|
||||||
// 3. Save to database
|
// 3. Save to database
|
||||||
db.createWorkspace(channel.id, guild.id, name, workspacePath);
|
db.createWorkspace(channel.id, guild.id, name, workspacePath);
|
||||||
|
|
||||||
// 4. Reply with confirmation
|
// 4. Reply with confirmation
|
||||||
await interaction.editReply(
|
await interaction.editReply({
|
||||||
`Agent **${name}** erstellt in <#${channel.id}>.\n` +
|
content:
|
||||||
|
`Agent **${name}** erstellt in <#${channel.id}>.\n` +
|
||||||
`Rolle: ${role}\n` +
|
`Rolle: ${role}\n` +
|
||||||
`Workspace: \`${workspacePath}\``
|
`Modell: ${selectedModel}\n` +
|
||||||
);
|
`Workspace: \`${workspacePath}\``,
|
||||||
|
components: [],
|
||||||
|
});
|
||||||
|
|
||||||
// 5. Send welcome message in the new channel
|
// 5. Send welcome message in the new channel
|
||||||
await channel.send(
|
await channel.send(
|
||||||
|
|
@ -115,6 +167,9 @@ export async function handleNewAgent(
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
const msg =
|
const msg =
|
||||||
error instanceof Error ? error.message : "Unbekannter Fehler";
|
error instanceof Error ? error.message : "Unbekannter Fehler";
|
||||||
await interaction.editReply(`Fehler beim Erstellen des Agenten: ${msg}`);
|
await interaction.editReply({
|
||||||
|
content: `Fehler beim Erstellen des Agenten: ${msg}`,
|
||||||
|
components: [],
|
||||||
|
});
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
8
src/config/models.ts
Normal file
8
src/config/models.ts
Normal file
|
|
@ -0,0 +1,8 @@
|
||||||
|
export const ALLOWED_MODELS = [
|
||||||
|
{ id: "claude-opus-4-6", label: "Claude Opus 4.6 (leistungsstark, langsamer)" },
|
||||||
|
{ id: "claude-sonnet-4-6", label: "Claude Sonnet 4.6 (Standard, empfohlen)" },
|
||||||
|
{ id: "claude-haiku-4-5-20251001", label: "Claude Haiku 4.5 (schnell, günstig)" },
|
||||||
|
] as const;
|
||||||
|
|
||||||
|
export type AllowedModel = typeof ALLOWED_MODELS[number]["id"];
|
||||||
|
export const DEFAULT_MODEL: AllowedModel = "claude-sonnet-4-6";
|
||||||
|
|
@ -1,12 +1,12 @@
|
||||||
---
|
---
|
||||||
id: DIS-153
|
id: DIS-153
|
||||||
status: ready
|
status: in-progress
|
||||||
phase: 1.5
|
phase: 1.5
|
||||||
priority: p2
|
priority: p2
|
||||||
labels: [phase:1.5, type:feat, priority:p2]
|
labels: [phase:1.5, type:feat, priority:p2]
|
||||||
branch: refinement/model-selection
|
branch: refinement/model-selection
|
||||||
assignee: null
|
assignee: developer-agent
|
||||||
started: null
|
started: 2026-04-13
|
||||||
pr: null
|
pr: null
|
||||||
merged: null
|
merged: null
|
||||||
---
|
---
|
||||||
|
|
|
||||||
165
tests/unit/runner-model.test.ts
Normal file
165
tests/unit/runner-model.test.ts
Normal file
|
|
@ -0,0 +1,165 @@
|
||||||
|
import { describe, it, expect, vi, beforeEach, afterEach, beforeAll, afterAll } from "vitest";
|
||||||
|
import * as fs from "fs";
|
||||||
|
import * as os from "os";
|
||||||
|
import * as path from "path";
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Unit tests for DIS-153: model selection per agent.
|
||||||
|
*
|
||||||
|
* cross-spawn is mocked so no real process is launched.
|
||||||
|
* Each test inspects the args passed to spawn to verify --model is set correctly.
|
||||||
|
*/
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Spy on the args cross-spawn receives
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
let capturedArgs: string[] = [];
|
||||||
|
|
||||||
|
vi.mock("cross-spawn", () => {
|
||||||
|
function makeMockChild(args: string[]) {
|
||||||
|
capturedArgs = args;
|
||||||
|
|
||||||
|
const stdoutListeners: Record<string, Array<(...a: unknown[]) => void>> = {};
|
||||||
|
|
||||||
|
const stdoutStream = {
|
||||||
|
setEncoding: () => {},
|
||||||
|
on: (event: string, cb: (...a: unknown[]) => void) => {
|
||||||
|
if (!stdoutListeners[event]) stdoutListeners[event] = [];
|
||||||
|
stdoutListeners[event].push(cb);
|
||||||
|
|
||||||
|
if (event === "data") {
|
||||||
|
Promise.resolve().then(() => {
|
||||||
|
for (const handler of stdoutListeners["data"] ?? []) {
|
||||||
|
handler(JSON.stringify({ result: "ok" }));
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
const listeners: Record<string, Array<(...a: unknown[]) => void>> = {};
|
||||||
|
const emit = (event: string, ...a: unknown[]) => {
|
||||||
|
for (const cb of listeners[event] ?? []) cb(...a);
|
||||||
|
};
|
||||||
|
|
||||||
|
const child = {
|
||||||
|
stdout: stdoutStream,
|
||||||
|
stderr: { setEncoding: () => {}, on: () => {} },
|
||||||
|
stdin: { end: () => {} },
|
||||||
|
on: (event: string, cb: (...a: unknown[]) => void) => {
|
||||||
|
if (!listeners[event]) listeners[event] = [];
|
||||||
|
listeners[event].push(cb);
|
||||||
|
|
||||||
|
if (event === "close") {
|
||||||
|
Promise.resolve()
|
||||||
|
.then(() => Promise.resolve())
|
||||||
|
.then(() => emit("close", 0));
|
||||||
|
}
|
||||||
|
},
|
||||||
|
kill: () => {},
|
||||||
|
};
|
||||||
|
|
||||||
|
return child;
|
||||||
|
}
|
||||||
|
|
||||||
|
return {
|
||||||
|
default: (_cmd: string, args: string[]) => makeMockChild(args),
|
||||||
|
};
|
||||||
|
});
|
||||||
|
|
||||||
|
import { runAgent, _resetSemaphore } from "../../src/agent/runner";
|
||||||
|
import { _resetResolveClaudeCache } from "../../src/runtime/resolve-claude";
|
||||||
|
import { AgentYamlSchema } from "../../src/agent/schema";
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Shared workspace fixture
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
describe("runner model selection (DIS-153)", () => {
|
||||||
|
let workspaceDir: string;
|
||||||
|
const originalEnv = { ...process.env };
|
||||||
|
|
||||||
|
beforeAll(() => {
|
||||||
|
workspaceDir = fs.mkdtempSync(path.join(os.tmpdir(), "disclaw-model-test-"));
|
||||||
|
fs.writeFileSync(
|
||||||
|
path.join(workspaceDir, "CLAUDE.md"),
|
||||||
|
"# Model Test Agent\nYou are a test agent.\n",
|
||||||
|
"utf-8"
|
||||||
|
);
|
||||||
|
});
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
capturedArgs = [];
|
||||||
|
_resetResolveClaudeCache();
|
||||||
|
_resetSemaphore(4);
|
||||||
|
process.env.CLAUDE_PATH = "/fake/claude";
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
process.env = { ...originalEnv };
|
||||||
|
_resetResolveClaudeCache();
|
||||||
|
});
|
||||||
|
|
||||||
|
afterAll(() => {
|
||||||
|
try {
|
||||||
|
fs.rmSync(workspaceDir, { recursive: true, force: true });
|
||||||
|
} catch {
|
||||||
|
// best-effort cleanup
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
// -------------------------------------------------------------------------
|
||||||
|
// Test 1: --model is passed when model is set in agent.yaml
|
||||||
|
// -------------------------------------------------------------------------
|
||||||
|
it("passes --model claude-sonnet-4-6 when model is set in agent.yaml", async () => {
|
||||||
|
fs.writeFileSync(
|
||||||
|
path.join(workspaceDir, "agent.yaml"),
|
||||||
|
`name: model-agent\ndisplay_name: Model Agent\nrole: Testing\nchannel_id: "111"\nmodel: claude-sonnet-4-6\n`,
|
||||||
|
"utf-8"
|
||||||
|
);
|
||||||
|
|
||||||
|
await runAgent({
|
||||||
|
workspacePath: workspaceDir,
|
||||||
|
channelName: "test",
|
||||||
|
userMessage: "hello",
|
||||||
|
conversationHistory: [],
|
||||||
|
});
|
||||||
|
|
||||||
|
const modelIdx = capturedArgs.indexOf("--model");
|
||||||
|
expect(modelIdx).toBeGreaterThan(-1);
|
||||||
|
expect(capturedArgs[modelIdx + 1]).toBe("claude-sonnet-4-6");
|
||||||
|
});
|
||||||
|
|
||||||
|
// -------------------------------------------------------------------------
|
||||||
|
// Test 2: Default model is used when model field is absent from agent.yaml
|
||||||
|
// -------------------------------------------------------------------------
|
||||||
|
it("passes default model when model field is absent from agent.yaml", async () => {
|
||||||
|
fs.writeFileSync(
|
||||||
|
path.join(workspaceDir, "agent.yaml"),
|
||||||
|
`name: model-agent\ndisplay_name: Model Agent\nrole: Testing\nchannel_id: "222"\n`,
|
||||||
|
"utf-8"
|
||||||
|
);
|
||||||
|
|
||||||
|
await runAgent({
|
||||||
|
workspacePath: workspaceDir,
|
||||||
|
channelName: "test",
|
||||||
|
userMessage: "hello",
|
||||||
|
conversationHistory: [],
|
||||||
|
});
|
||||||
|
|
||||||
|
const modelIdx = capturedArgs.indexOf("--model");
|
||||||
|
expect(modelIdx).toBeGreaterThan(-1);
|
||||||
|
// DEFAULT_MODEL is claude-sonnet-4-6
|
||||||
|
expect(capturedArgs[modelIdx + 1]).toBe("claude-sonnet-4-6");
|
||||||
|
});
|
||||||
|
|
||||||
|
// -------------------------------------------------------------------------
|
||||||
|
// Test 3: Invalid model in agent.yaml → Zod validation throws
|
||||||
|
// -------------------------------------------------------------------------
|
||||||
|
it("throws when agent.yaml contains an invalid model value", () => {
|
||||||
|
const rawYaml = `name: bad-model-agent\ndisplay_name: Bad\nrole: Testing\nchannel_id: "333"\nmodel: gpt-4-turbo\n`;
|
||||||
|
const parsed = { name: "bad-model-agent", display_name: "Bad", role: "Testing", channel_id: "333", model: "gpt-4-turbo" };
|
||||||
|
|
||||||
|
expect(() => AgentYamlSchema.parse(parsed)).toThrow();
|
||||||
|
});
|
||||||
|
});
|
||||||
Loading…
Reference in a new issue