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 OpenAI from "openai";
|
||||||
|
import type { z } from "zod";
|
||||||
|
|
||||||
export interface GenerateInput {
|
export interface GenerateInput {
|
||||||
system?: string;
|
system?: string;
|
||||||
@@ -7,6 +8,10 @@ export interface GenerateInput {
|
|||||||
temperature?: number;
|
temperature?: number;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export interface GenerateValidatedJsonInput<T> extends GenerateInput {
|
||||||
|
schema: z.ZodType<T>;
|
||||||
|
}
|
||||||
|
|
||||||
export interface LlmProviderStatus {
|
export interface LlmProviderStatus {
|
||||||
provider: "deepseek" | "openai";
|
provider: "deepseek" | "openai";
|
||||||
configured: boolean;
|
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) {
|
function normalizeLlmError(error: unknown) {
|
||||||
const message = error instanceof Error ? error.message : "Unknown LLM error";
|
const message = error instanceof Error ? error.message : "Unknown LLM error";
|
||||||
return new Error(`LLM provider error: ${message}`);
|
return new Error(`LLM provider error: ${message}`);
|
||||||
|
|||||||
Reference in New Issue
Block a user