feat(agent): route chat through selected model
This commit is contained in:
@@ -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({
|
||||||
|
|||||||
@@ -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,
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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),
|
||||||
});
|
});
|
||||||
|
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
@@ -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;
|
||||||
|
});
|
||||||
|
|||||||
@@ -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);
|
||||||
|
});
|
||||||
});
|
});
|
||||||
|
|||||||
Reference in New Issue
Block a user