UNPKG

ai-sdk-guardrails

Version:

Input and output guardrails middleware for Vercel AI SDK.

608 lines (606 loc) 19.6 kB
import { createInputGuardrail, createOutputGuardrail } from "./chunk-HHQ3CIFN.js"; import { GuardrailConfigurationError, GuardrailExecutionError, GuardrailTimeoutError, GuardrailValidationError, GuardrailsError, InputBlockedError, MiddlewareError, OutputBlockedError, extractErrorInfo, isGuardrailsError } from "./chunk-LLCOPUS6.js"; // src/guardrails.ts import { wrapLanguageModel } from "ai"; function defineInputGuardrail(guardrail) { const enhanced = { enabled: true, priority: "medium", version: "1.0.0", tags: [], ...guardrail, execute: async (params) => { const startTime = Date.now(); const originalExecute = guardrail.execute; try { const result = await originalExecute(params); const executionTime = Date.now() - startTime; return { ...result, context: { guardrailName: guardrail.name, guardrailVersion: guardrail.version, executedAt: /* @__PURE__ */ new Date(), executionTimeMs: executionTime, ...result.context } }; } catch (error) { const executionTime = Date.now() - startTime; return { tripwireTriggered: true, message: `Guardrail execution failed: ${error instanceof Error ? error.message : "Unknown error"}`, severity: "critical", context: { guardrailName: guardrail.name, guardrailVersion: guardrail.version, executedAt: /* @__PURE__ */ new Date(), executionTimeMs: executionTime }, metadata: { error: error instanceof Error ? error.message : String(error) } }; } } }; return enhanced; } async function executeInputGuardrails(guardrails, params, options = {}) { const { parallel = true, timeout = 3e4, // 30 seconds continueOnFailure = true, logLevel = "info" } = options; const enabledGuardrails = guardrails.filter((g) => g.enabled !== false).sort((a, b) => { const priorityOrder = { critical: 4, high: 3, medium: 2, low: 1 }; return (priorityOrder[b.priority || "medium"] || 2) - (priorityOrder[a.priority || "medium"] || 2); }); const results = []; const executeWithTimeout = async (guardrail) => { const timeoutPromise = new Promise((_, reject) => { setTimeout(async () => { const { GuardrailTimeoutError: GuardrailTimeoutError2 } = await import("./errors-BTTWMQEI.js"); reject(new GuardrailTimeoutError2(guardrail.name, timeout)); }, timeout); }); const executionPromise = guardrail.execute(params); return Promise.race([executionPromise, timeoutPromise]); }; if (parallel) { const promises = enabledGuardrails.map(async (guardrail) => { try { const result = await executeWithTimeout(guardrail); if (result.tripwireTriggered && logLevel !== "none") { console.log( `Input guardrail "${guardrail.name}" triggered: ${result.message}` ); } return result; } catch (error) { if (logLevel !== "none") { console.error( `Error executing input guardrail "${guardrail.name}":`, error ); } return { tripwireTriggered: true, message: `Guardrail execution failed: ${error instanceof Error ? error.message : "Unknown error"}`, severity: "critical", metadata: { error: error instanceof Error ? error.message : String(error) } }; } }); results.push(...await Promise.all(promises)); } else { for (const guardrail of enabledGuardrails) { try { const result = await executeWithTimeout(guardrail); results.push(result); if (result.tripwireTriggered) { if (logLevel !== "none") { console.log( `Input guardrail "${guardrail.name}" triggered: ${result.message}` ); } if (!continueOnFailure) { break; } } } catch (error) { if (logLevel !== "none") { console.error( `Error executing input guardrail "${guardrail.name}":`, error ); } const errorResult = { tripwireTriggered: true, message: `Guardrail execution failed: ${error instanceof Error ? error.message : "Unknown error"}`, severity: "critical", metadata: { error: error instanceof Error ? error.message : String(error) } }; results.push(errorResult); if (!continueOnFailure) { break; } } } } return results; } function defineOutputGuardrail(guardrail) { const enhanced = { enabled: true, priority: "medium", version: "1.0.0", tags: [], ...guardrail, execute: async (params) => { const startTime = Date.now(); const originalExecute = guardrail.execute; try { const result = await originalExecute(params); const executionTime = Date.now() - startTime; return { ...result, context: { guardrailName: guardrail.name, guardrailVersion: guardrail.version, executedAt: /* @__PURE__ */ new Date(), executionTimeMs: executionTime, ...result.context } }; } catch (error) { const executionTime = Date.now() - startTime; return { tripwireTriggered: true, message: `Guardrail execution failed: ${error instanceof Error ? error.message : "Unknown error"}`, severity: "critical", context: { guardrailName: guardrail.name, guardrailVersion: guardrail.version, executedAt: /* @__PURE__ */ new Date(), executionTimeMs: executionTime }, metadata: { error: error instanceof Error ? error.message : String(error) } }; } } }; return enhanced; } async function executeOutputGuardrails(guardrails, params, options = {}) { const { parallel = true, timeout = 3e4, // 30 seconds continueOnFailure = true, logLevel = "info" } = options; const enabledGuardrails = guardrails.filter((g) => g.enabled !== false).sort((a, b) => { const priorityOrder = { critical: 4, high: 3, medium: 2, low: 1 }; return (priorityOrder[b.priority || "medium"] || 2) - (priorityOrder[a.priority || "medium"] || 2); }); const results = []; const executeWithTimeout = async (guardrail) => { const timeoutPromise = new Promise((_, reject) => { setTimeout(async () => { const { GuardrailTimeoutError: GuardrailTimeoutError2 } = await import("./errors-BTTWMQEI.js"); reject(new GuardrailTimeoutError2(guardrail.name, timeout)); }, timeout); }); const executionPromise = guardrail.execute(params); return Promise.race([executionPromise, timeoutPromise]); }; if (parallel) { const promises = enabledGuardrails.map(async (guardrail) => { try { const result = await executeWithTimeout(guardrail); if (result.tripwireTriggered && logLevel !== "none") { console.log( `Output guardrail "${guardrail.name}" triggered: ${result.message}` ); } return result; } catch (error) { if (logLevel !== "none") { console.error( `Error executing output guardrail "${guardrail.name}":`, error ); } return { tripwireTriggered: true, message: `Guardrail execution failed: ${error instanceof Error ? error.message : "Unknown error"}`, severity: "critical", metadata: { error: error instanceof Error ? error.message : String(error) } }; } }); results.push(...await Promise.all(promises)); } else { for (const guardrail of enabledGuardrails) { try { const result = await executeWithTimeout(guardrail); results.push(result); if (result.tripwireTriggered) { if (logLevel !== "none") { console.log( `Output guardrail "${guardrail.name}" triggered: ${result.message}` ); } if (!continueOnFailure) { break; } } } catch (error) { if (logLevel !== "none") { console.error( `Error executing output guardrail "${guardrail.name}":`, error ); } const errorResult = { tripwireTriggered: true, message: `Guardrail execution failed: ${error instanceof Error ? error.message : "Unknown error"}`, severity: "critical", metadata: { error: error instanceof Error ? error.message : String(error) } }; results.push(errorResult); if (!continueOnFailure) { break; } } } } return results; } function extractTextFromContent(content) { return content.filter((part) => part.type === "text").map((part) => part.text || "").join(""); } function wrapWithInputGuardrails(model, guardrails, options) { const middleware = createInputGuardrailsMiddleware({ inputGuardrails: guardrails, ...options }); return wrapLanguageModel({ model, middleware }); } function wrapWithOutputGuardrails(model, guardrails, options) { const middleware = createOutputGuardrailsMiddleware({ outputGuardrails: guardrails, ...options }); return wrapLanguageModel({ model, middleware }); } function wrapWithGuardrails(model, config) { const { inputGuardrails = [], outputGuardrails = [], throwOnBlocked, executionOptions, onInputBlocked, onOutputBlocked } = config; const middlewares = []; if (inputGuardrails.length > 0) { middlewares.push( createInputGuardrailsMiddleware({ inputGuardrails, throwOnBlocked, executionOptions, onInputBlocked }) ); } if (outputGuardrails.length > 0) { middlewares.push( createOutputGuardrailsMiddleware({ outputGuardrails, throwOnBlocked, executionOptions, onOutputBlocked }) ); } if (middlewares.length === 0) { return model; } return wrapLanguageModel({ model, middleware: middlewares }); } function createInputGuardrailsMiddleware(config) { const { inputGuardrails, executionOptions = {}, onInputBlocked, throwOnBlocked = false } = config; return { transformParams: async ({ params }) => { const enhancedParams = { ...params, guardrailsBlocked: void 0 }; const promptMessages = Array.isArray(enhancedParams.prompt) ? enhancedParams.prompt : []; const systemMessage = promptMessages.find((msg) => msg.role === "system"); const system = systemMessage && Array.isArray(systemMessage.content) ? extractTextFromContent(systemMessage.content) : ""; const messages = promptMessages.filter((msg) => msg.role !== "system").map((msg) => ({ role: msg.role, content: msg.content && Array.isArray(msg.content) ? extractTextFromContent(msg.content) : "" })); const prompt = messages.length === 1 && messages[0]?.role === "user" ? messages[0].content : messages.map((m) => m.content).join(" "); const guardrailContext = { prompt, messages, system, maxOutputTokens: enhancedParams.maxOutputTokens, temperature: enhancedParams.temperature }; const inputResults = await executeInputGuardrails( inputGuardrails, guardrailContext, executionOptions ); const blockedResults = inputResults.filter((r) => r.tripwireTriggered); if (blockedResults.length > 0) { if (onInputBlocked) { onInputBlocked(blockedResults, guardrailContext); } if (throwOnBlocked) { const { InputBlockedError: InputBlockedError2 } = await import("./errors-BTTWMQEI.js"); const blockedGuardrails = blockedResults.map((r) => ({ name: r.context?.guardrailName || "unknown", message: r.message || "Blocked", severity: r.severity || "medium" })); throw new InputBlockedError2(blockedGuardrails); } enhancedParams.guardrailsBlocked = blockedResults; } return enhancedParams; }, wrapGenerate: async ({ doGenerate, params }) => { const paramsWithGuardrails = params; if (paramsWithGuardrails.guardrailsBlocked) { const blockedResults = paramsWithGuardrails.guardrailsBlocked; const blockedMessage = blockedResults.map((r) => r.message).join(", "); return { content: [ { type: "text", text: `[Input blocked: ${blockedMessage}]` } ], finishReason: "other", usage: { inputTokens: 0, outputTokens: 0, totalTokens: 0 }, warnings: [] }; } return doGenerate(); }, wrapStream: async ({ doStream, params }) => { const paramsWithGuardrails = params; if (paramsWithGuardrails.guardrailsBlocked) { const blockedResults = paramsWithGuardrails.guardrailsBlocked; const blockedMessage = blockedResults.map((r) => r.message).join(", "); const stream = new ReadableStream({ start(controller) { controller.enqueue({ type: "text-delta", id: "1", delta: `[Input blocked: ${blockedMessage}]` }); controller.enqueue({ type: "finish", finishReason: "other", usage: { inputTokens: 0, outputTokens: 0, totalTokens: 0 } }); controller.close(); } }); return { stream }; } return doStream(); } }; } function createOutputGuardrailsMiddleware(config) { const { outputGuardrails, executionOptions = {}, onOutputBlocked, throwOnBlocked = false } = config; return { wrapGenerate: async ({ doGenerate, params }) => { const result = await doGenerate(); const promptMessages = Array.isArray(params.prompt) ? params.prompt : []; const systemMessage = promptMessages.find((msg) => msg.role === "system"); const system = systemMessage && Array.isArray(systemMessage.content) ? extractTextFromContent(systemMessage.content) : ""; const messages = promptMessages.filter((msg) => msg.role !== "system").map((msg) => ({ role: msg.role, content: msg.content && Array.isArray(msg.content) ? extractTextFromContent(msg.content) : "" })); const prompt = messages.length === 1 && messages[0]?.role === "user" ? messages[0].content : messages.map((m) => m.content).join(" "); const guardrailContext = { prompt, messages, system, maxOutputTokens: params.maxOutputTokens, temperature: params.temperature }; const outputContext = { input: guardrailContext, result }; const outputResults = await executeOutputGuardrails( outputGuardrails, outputContext, executionOptions ); const blockedResults = outputResults.filter((r) => r.tripwireTriggered); if (blockedResults.length > 0) { if (onOutputBlocked) { onOutputBlocked(blockedResults, guardrailContext, result); } if (throwOnBlocked) { const { OutputBlockedError: OutputBlockedError2 } = await import("./errors-BTTWMQEI.js"); const blockedGuardrails = blockedResults.map((r) => ({ name: r.context?.guardrailName || "unknown", message: r.message || "Blocked", severity: r.severity || "medium" })); throw new OutputBlockedError2(blockedGuardrails); } } return result; }, wrapStream: async ({ doStream, params }) => { const streamResult = await doStream(); let accumulatedText = ""; const blockedChunks = []; const transformStream = new TransformStream({ transform(chunk) { if (chunk.type === "text-delta") { accumulatedText += chunk.delta || ""; } blockedChunks.push(chunk); }, async flush(controller) { const promptMessages = Array.isArray(params.prompt) ? params.prompt : []; const systemMessage = promptMessages.find( (msg) => msg.role === "system" ); const system = systemMessage && Array.isArray(systemMessage.content) ? extractTextFromContent(systemMessage.content) : ""; const messages = promptMessages.filter((msg) => msg.role !== "system").map((msg) => ({ role: msg.role, content: msg.content && Array.isArray(msg.content) ? extractTextFromContent(msg.content) : "" })); const prompt = messages.length === 1 && messages[0]?.role === "user" ? messages[0].content : messages.map((m) => m.content).join(" "); const guardrailContext = { prompt, messages, system, maxOutputTokens: params.maxOutputTokens, temperature: params.temperature }; const outputContext = { input: guardrailContext, result: { text: accumulatedText } }; const outputResults = await executeOutputGuardrails( outputGuardrails, outputContext, executionOptions ); const blockedResults = outputResults.filter( (r) => r.tripwireTriggered ); if (blockedResults.length > 0) { if (onOutputBlocked) { onOutputBlocked(blockedResults, guardrailContext, { text: accumulatedText }); } if (throwOnBlocked) { controller.error( new Error( `Output guardrails blocked response: ${blockedResults.map((r) => r.message).join(", ")}` ) ); return; } controller.enqueue({ type: "text-delta", id: "1", delta: "[Output blocked by guardrails]" }); controller.enqueue({ type: "finish", finishReason: "error", usage: { inputTokens: 0, outputTokens: 0, totalTokens: 0 } }); } else { for (const chunk of blockedChunks) { controller.enqueue(chunk); } } } }); return { stream: streamResult.stream.pipeThrough(transformStream) }; } }; } export { GuardrailConfigurationError, GuardrailExecutionError, GuardrailTimeoutError, GuardrailValidationError, GuardrailsError, InputBlockedError, MiddlewareError, OutputBlockedError, createInputGuardrail, createInputGuardrailsMiddleware, createOutputGuardrail, createOutputGuardrailsMiddleware, defineInputGuardrail, defineOutputGuardrail, executeInputGuardrails, executeOutputGuardrails, extractErrorInfo, isGuardrailsError, wrapWithGuardrails, wrapWithInputGuardrails, wrapWithOutputGuardrails };