ai-sdk-guardrails
Version:
Input and output guardrails middleware for Vercel AI SDK.
608 lines (606 loc) • 19.6 kB
JavaScript
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
};