feat(agent): route chat through selected model

This commit is contained in:
ginnoir
2026-07-08 16:34:22 -05:00
parent bf7e07ead9
commit 9b7a04431c
6 changed files with 78 additions and 4 deletions
+15
View File
@@ -2,6 +2,7 @@ import { apiError, apiJson } from "@/lib/api-handler";
import { resolveApiAuth } from "@/lib/api-auth"; import { resolveApiAuth } from "@/lib/api-auth";
import { getAssistantPreferences, resolveAssistantSystemPrompt } from "@/lib/assistant-preference"; import { getAssistantPreferences, resolveAssistantSystemPrompt } from "@/lib/assistant-preference";
import { isLlmConfigured } from "@/lib/llm"; import { isLlmConfigured } from "@/lib/llm";
import { listLlmModels, resolveAssistantModel } from "@/lib/llm/models";
import { clientChatInputSchema } from "@/modules/agent/messages"; import { clientChatInputSchema } from "@/modules/agent/messages";
import { encodeSseEvent } from "@/modules/agent/server/progress"; import { encodeSseEvent } from "@/modules/agent/server/progress";
import { runAgentChat } from "@/modules/agent/server/run"; import { runAgentChat } from "@/modules/agent/server/run";
@@ -31,6 +32,18 @@ export async function POST(request: Request) {
return apiError(parsed.error.issues[0]?.message ?? "Validation error", 400); return apiError(parsed.error.issues[0]?.message ?? "Validation error", 400);
} }
const modelList = await listLlmModels();
const modelResolution = resolveAssistantModel({
requestedModel: parsed.data.model,
savedModel: assistant.model,
fallbackModel: modelList.fallbackModel,
models: modelList.models,
});
if (!modelResolution.ok) {
return apiError(modelResolution.error, 400);
}
if (parsed.data.stream) { if (parsed.data.stream) {
const stream = new ReadableStream<Uint8Array>({ const stream = new ReadableStream<Uint8Array>({
async start(controller) { async start(controller) {
@@ -44,6 +57,7 @@ export async function POST(request: Request) {
messages: parsed.data.messages, messages: parsed.data.messages,
request, request,
systemPrompt, systemPrompt,
model: modelResolution.model,
onProgress: send, onProgress: send,
}); });
@@ -78,6 +92,7 @@ export async function POST(request: Request) {
messages: parsed.data.messages, messages: parsed.data.messages,
request, request,
systemPrompt, systemPrompt,
model: modelResolution.model,
}); });
return apiJson({ return apiJson({
+3 -3
View File
@@ -13,8 +13,8 @@ export type {
export { getLlmConfig, isLlmConfigured } from "./config"; export { getLlmConfig, isLlmConfigured } from "./config";
export { createMockLlmClient } from "./mock"; export { createMockLlmClient } from "./mock";
export function createLlmClient(override?: LlmClient): LlmClient { export function createLlmClient(options?: { model?: string; override?: LlmClient }): LlmClient {
if (override) return override; if (options?.override) return options.override;
const config = getLlmConfig(); const config = getLlmConfig();
if (config.provider === "mock" || !config.baseUrl) { if (config.provider === "mock" || !config.baseUrl) {
@@ -24,6 +24,6 @@ export function createLlmClient(override?: LlmClient): LlmClient {
return createOpenAiCompatibleClient({ return createOpenAiCompatibleClient({
baseUrl: config.baseUrl, baseUrl: config.baseUrl,
apiKey: config.apiKey, apiKey: config.apiKey,
model: config.model, model: options?.model ?? config.model,
}); });
} }
+7
View File
@@ -1,4 +1,5 @@
import { z } from "zod"; import { z } from "zod";
import { isValidLlmModelId } from "@/lib/llm/models";
export const clientChatAttachmentSchema = z.object({ export const clientChatAttachmentSchema = z.object({
type: z.literal("image"), type: z.literal("image"),
@@ -11,8 +12,14 @@ export const clientChatMessageSchema = z.object({
attachments: z.array(clientChatAttachmentSchema).max(3).optional(), attachments: z.array(clientChatAttachmentSchema).max(3).optional(),
}); });
export const clientChatModelSchema = z
.string()
.trim()
.refine((value) => isValidLlmModelId(value), "Invalid assistant model");
export const clientChatInputSchema = z.object({ export const clientChatInputSchema = z.object({
stream: z.boolean().optional(), stream: z.boolean().optional(),
model: clientChatModelSchema.optional(),
messages: z.array(clientChatMessageSchema).min(1).max(40), messages: z.array(clientChatMessageSchema).min(1).max(40),
}); });
+2 -1
View File
@@ -45,11 +45,12 @@ export async function runAgentChat(options: {
messages: ClientChatMessage[]; messages: ClientChatMessage[];
request: Request; request: Request;
systemPrompt?: string; systemPrompt?: string;
model?: string;
llm?: LlmClient; llm?: LlmClient;
executeTool?: ToolExecutor; executeTool?: ToolExecutor;
onProgress?: AgentProgressHandler; onProgress?: AgentProgressHandler;
}): Promise<AgentChatResult> { }): Promise<AgentChatResult> {
const llm = options.llm ?? createLlmClient(); const llm = options.llm ?? createLlmClient({ model: options.model });
const executeTool = options.executeTool ?? createApiToolExecutor(options.request); const executeTool = options.executeTool ?? createApiToolExecutor(options.request);
const onProgress = options.onProgress; const onProgress = options.onProgress;
const systemPrompt = options.systemPrompt ?? AGENT_SYSTEM_PROMPT; const systemPrompt = options.systemPrompt ?? AGENT_SYSTEM_PROMPT;
+33
View File
@@ -84,3 +84,36 @@ describe("runAgentChat", () => {
assert.ok(result.message.content.length > 0); assert.ok(result.message.content.length > 0);
}); });
}); });
it("passes a model override to the OpenAI-compatible client", async () => {
const originalBaseUrl = process.env.LLM_BASE_URL;
const originalModel = process.env.LLM_MODEL;
const originalProvider = process.env.LLM_PROVIDER;
const originalFetch = globalThis.fetch;
let requestBody: unknown = null;
process.env.LLM_BASE_URL = "https://llm.example.test/v1";
process.env.LLM_MODEL = "llama3.2";
delete process.env.LLM_PROVIDER;
globalThis.fetch = (async (_input: RequestInfo | URL, init?: RequestInit) => {
requestBody = JSON.parse(String(init?.body));
return Response.json({
choices: [{ message: { role: "assistant", content: "done" }, finish_reason: "stop" }],
});
}) as typeof fetch;
const { createLlmClient } = await import("../../src/lib/llm/index");
const client = createLlmClient({ model: "qwen2.5-coder" });
await client.chatCompletion({ messages: [{ role: "user", content: "hello" }] });
assert.equal((requestBody as { model?: string }).model, "qwen2.5-coder");
globalThis.fetch = originalFetch;
if (originalBaseUrl === undefined) delete process.env.LLM_BASE_URL;
else process.env.LLM_BASE_URL = originalBaseUrl;
if (originalModel === undefined) delete process.env.LLM_MODEL;
else process.env.LLM_MODEL = originalModel;
if (originalProvider === undefined) delete process.env.LLM_PROVIDER;
else process.env.LLM_PROVIDER = originalProvider;
});
+18
View File
@@ -25,4 +25,22 @@ describe("clientChatInputSchema", () => {
assert.equal(parsed.success, false); assert.equal(parsed.success, false);
}); });
it("accepts an optional model ID", () => {
const parsed = clientChatInputSchema.safeParse({
model: "qwen2.5-coder",
messages: [{ role: "user", content: "hello" }],
});
assert.equal(parsed.success, true);
});
it("rejects invalid model IDs", () => {
const parsed = clientChatInputSchema.safeParse({
model: "bad model",
messages: [{ role: "user", content: "hello" }],
});
assert.equal(parsed.success, false);
});
}); });