feat: add validated llm json helper
This commit is contained in:
@@ -0,0 +1,69 @@
|
||||
import { afterEach, describe, expect, it, vi } from "vitest";
|
||||
import { z } from "zod";
|
||||
|
||||
import * as client from "../client";
|
||||
|
||||
describe("generateValidatedJson", () => {
|
||||
const originalProvider = process.env.LLM_PROVIDER;
|
||||
const originalDeepSeekKey = process.env.DEEPSEEK_API_KEY;
|
||||
|
||||
afterEach(() => {
|
||||
process.env.LLM_PROVIDER = originalProvider;
|
||||
process.env.DEEPSEEK_API_KEY = originalDeepSeekKey;
|
||||
client.setGenerateJsonForValidation(client.generateJson);
|
||||
vi.restoreAllMocks();
|
||||
});
|
||||
|
||||
it("returns null when no provider key is configured", async () => {
|
||||
process.env.LLM_PROVIDER = "deepseek";
|
||||
delete process.env.DEEPSEEK_API_KEY;
|
||||
|
||||
const result = await client.generateValidatedJson({
|
||||
schema: z.object({ value: z.string() }),
|
||||
prompt: "Return JSON.",
|
||||
});
|
||||
|
||||
expect(result).toBeNull();
|
||||
});
|
||||
|
||||
it("returns parsed data when the model response matches the schema", async () => {
|
||||
process.env.LLM_PROVIDER = "deepseek";
|
||||
process.env.DEEPSEEK_API_KEY = "test-key";
|
||||
client.setGenerateJsonForValidation(async () => ({ value: "from-llm" }));
|
||||
|
||||
const result = await client.generateValidatedJson({
|
||||
schema: z.object({ value: z.string() }),
|
||||
prompt: "Return JSON.",
|
||||
});
|
||||
|
||||
expect(result).toEqual({ value: "from-llm" });
|
||||
});
|
||||
|
||||
it("returns null when the model response fails schema validation", async () => {
|
||||
process.env.LLM_PROVIDER = "deepseek";
|
||||
process.env.DEEPSEEK_API_KEY = "test-key";
|
||||
client.setGenerateJsonForValidation(async () => ({ value: 42 }));
|
||||
|
||||
const result = await client.generateValidatedJson({
|
||||
schema: z.object({ value: z.string() }),
|
||||
prompt: "Return JSON.",
|
||||
});
|
||||
|
||||
expect(result).toBeNull();
|
||||
});
|
||||
|
||||
it("returns null when the provider call rejects", async () => {
|
||||
process.env.LLM_PROVIDER = "deepseek";
|
||||
process.env.DEEPSEEK_API_KEY = "test-key";
|
||||
client.setGenerateJsonForValidation(async () => {
|
||||
throw new Error("provider down");
|
||||
});
|
||||
|
||||
const result = await client.generateValidatedJson({
|
||||
schema: z.object({ value: z.string() }),
|
||||
prompt: "Return JSON.",
|
||||
});
|
||||
|
||||
expect(result).toBeNull();
|
||||
});
|
||||
});
|
||||
@@ -1,4 +1,5 @@
|
||||
import OpenAI from "openai";
|
||||
import type { z } from "zod";
|
||||
|
||||
export interface GenerateInput {
|
||||
system?: string;
|
||||
@@ -7,6 +8,10 @@ export interface GenerateInput {
|
||||
temperature?: number;
|
||||
}
|
||||
|
||||
export interface GenerateValidatedJsonInput<T> extends GenerateInput {
|
||||
schema: z.ZodType<T>;
|
||||
}
|
||||
|
||||
export interface LlmProviderStatus {
|
||||
provider: "deepseek" | "openai";
|
||||
configured: boolean;
|
||||
@@ -106,6 +111,31 @@ export async function generateJson<T>(input: GenerateInput): Promise<T> {
|
||||
}
|
||||
}
|
||||
|
||||
export let generateJsonForValidation: <T>(input: GenerateInput) => Promise<T> =
|
||||
generateJson;
|
||||
|
||||
export function setGenerateJsonForValidation(
|
||||
generator: typeof generateJsonForValidation,
|
||||
) {
|
||||
generateJsonForValidation = generator;
|
||||
}
|
||||
|
||||
export async function generateValidatedJson<T>({
|
||||
schema,
|
||||
...input
|
||||
}: GenerateValidatedJsonInput<T>): Promise<T | null> {
|
||||
if (!isLlmConfigured()) {
|
||||
return null;
|
||||
}
|
||||
|
||||
try {
|
||||
const generated = await generateJsonForValidation<unknown>(input);
|
||||
return schema.parse(generated);
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
function normalizeLlmError(error: unknown) {
|
||||
const message = error instanceof Error ? error.message : "Unknown LLM error";
|
||||
return new Error(`LLM provider error: ${message}`);
|
||||
|
||||
Reference in New Issue
Block a user