From b9cfa3faf6f768e9c6738ce482a2fdb451365a35 Mon Sep 17 00:00:00 2001 From: Codex Date: Tue, 16 Jun 2026 22:46:45 +0800 Subject: [PATCH] feat: add validated llm json helper --- src/lib/llm/__tests__/client.test.ts | 69 ++++++++++++++++++++++++++++ src/lib/llm/client.ts | 30 ++++++++++++ 2 files changed, 99 insertions(+) create mode 100644 src/lib/llm/__tests__/client.test.ts diff --git a/src/lib/llm/__tests__/client.test.ts b/src/lib/llm/__tests__/client.test.ts new file mode 100644 index 0000000..f540435 --- /dev/null +++ b/src/lib/llm/__tests__/client.test.ts @@ -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(); + }); +}); diff --git a/src/lib/llm/client.ts b/src/lib/llm/client.ts index d6fc523..a3e7bea 100644 --- a/src/lib/llm/client.ts +++ b/src/lib/llm/client.ts @@ -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 extends GenerateInput { + schema: z.ZodType; +} + export interface LlmProviderStatus { provider: "deepseek" | "openai"; configured: boolean; @@ -106,6 +111,31 @@ export async function generateJson(input: GenerateInput): Promise { } } +export let generateJsonForValidation: (input: GenerateInput) => Promise = + generateJson; + +export function setGenerateJsonForValidation( + generator: typeof generateJsonForValidation, +) { + generateJsonForValidation = generator; +} + +export async function generateValidatedJson({ + schema, + ...input +}: GenerateValidatedJsonInput): Promise { + if (!isLlmConfigured()) { + return null; + } + + try { + const generated = await generateJsonForValidation(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}`);