ai-sdk-guardrails
Version:
Input and output guardrails middleware for Vercel AI SDK.
384 lines (380 loc) • 13.2 kB
JavaScript
import {
extractContent,
stringifyContent
} from "./chunk-WC2PXTRI.js";
import {
createOutputGuardrail
} from "./chunk-F7POYYOU.js";
// src/guardrails/prompt-leak.ts
var MAX_OUTPUT_LENGTH = 1024 * 1024;
var RE_NON_WORD = /[^\w\s]/g;
var RE_WHITESPACE = /\s+/;
var RE_REGEX_META = /[.*+?^${}()|[\]\\]/g;
function tokenize(text) {
return text.toLowerCase().replaceAll(RE_NON_WORD, " ").split(RE_WHITESPACE).filter(Boolean);
}
function generateNgrams(tokens, n) {
const ngrams = /* @__PURE__ */ new Set();
for (let i = 0; i <= tokens.length - n; i++) {
ngrams.add(tokens.slice(i, i + n).join(" "));
}
return ngrams;
}
function wordOverlapRatio(outputTokens, promptTokens) {
const outSet = new Set(outputTokens);
const promptSet = new Set(promptTokens);
let intersection = 0;
for (const w of outSet) {
if (promptSet.has(w)) intersection++;
}
const union = outSet.size + promptSet.size - intersection;
return union > 0 ? intersection / union : 0;
}
function findMatchingSubstrings(outputTokens, promptNgrams, ngramSize) {
const matches = [];
const windowSize = ngramSize + 4;
for (let i = 0; i <= outputTokens.length - ngramSize; i++) {
const ngram = outputTokens.slice(i, i + ngramSize).join(" ");
if (promptNgrams.has(ngram)) {
const end = Math.min(i + windowSize, outputTokens.length);
const fragment = outputTokens.slice(i, end).join(" ");
if (matches.every((m) => !m.includes(ngram))) {
matches.push(fragment);
}
}
}
return matches;
}
function detectSystemPromptLeak(output, systemPrompt, options = {}) {
if (!output || typeof output !== "string" || !systemPrompt || typeof systemPrompt !== "string") {
return {
leaked: false,
confidence: 0,
fragments: [],
sanitized: output || ""
};
}
const boundedOutput = output.length > MAX_OUTPUT_LENGTH ? output.slice(0, MAX_OUTPUT_LENGTH) : output;
const ngramSize = options.ngramSize ?? 4;
const threshold = options.threshold ?? 0.7;
const wordOverlapThreshold = options.wordOverlapThreshold ?? 0.25;
const redactionText = options.redactionText || "[REDACTED]";
const promptTokens = tokenize(systemPrompt);
const outputTokens = tokenize(boundedOutput);
if (promptTokens.length < 2) {
return {
leaked: false,
confidence: 0,
fragments: [],
sanitized: boundedOutput
};
}
const effectiveNgram = Math.min(ngramSize, Math.max(2, promptTokens.length));
const promptNgrams = generateNgrams(promptTokens, effectiveNgram);
const fragments = findMatchingSubstrings(
outputTokens,
promptNgrams,
effectiveNgram
);
const smallNgramSize = Math.min(3, Math.max(1, effectiveNgram - 1));
const smallFragments = smallNgramSize >= 2 && promptTokens.length >= smallNgramSize ? findMatchingSubstrings(
outputTokens,
generateNgrams(promptTokens, smallNgramSize),
smallNgramSize
) : [];
const outputNgrams = generateNgrams(outputTokens, effectiveNgram);
let ngramOverlap = 0;
for (const ng of outputNgrams) {
if (promptNgrams.has(ng)) ngramOverlap++;
}
const ngramOverlapRatio = promptNgrams.size > 0 ? ngramOverlap / promptNgrams.size : 0;
const wordOverlap = wordOverlapRatio(outputTokens, promptTokens);
const confidence = fragments.length > 0 ? Math.min(1, ngramOverlapRatio * 2 + (fragments.length > 2 ? 0.2 : 0)) : wordOverlap >= wordOverlapThreshold ? Math.min(1, wordOverlap * 2) : 0;
const isLeak = fragments.length > 0 && confidence >= threshold || fragments.length >= 2 || smallFragments.length >= 3 && wordOverlap >= wordOverlapThreshold || wordOverlap >= wordOverlapThreshold * 1.5 && smallFragments.length > 0;
if (!isLeak) {
return {
leaked: false,
confidence,
fragments: [],
sanitized: boundedOutput
};
}
const allFragments = [.../* @__PURE__ */ new Set([...fragments, ...smallFragments])];
let sanitized = boundedOutput;
for (const fragment of allFragments) {
const words = fragment.split(" ");
for (let len = words.length; len >= effectiveNgram; len--) {
const sub = words.slice(0, len).join(" ");
const regex = new RegExp(
sub.replaceAll(RE_REGEX_META, String.raw`\$&`).replaceAll(/\s+/g, String.raw`\s+`),
"gi"
);
sanitized = sanitized.replace(regex, () => redactionText);
}
}
return { leaked: true, confidence, fragments: allFragments, sanitized };
}
function systemPromptLeakDetector(options = {}) {
return createOutputGuardrail(
"system-prompt-leak-detector",
(context, accumulatedText) => {
const systemPrompt = options.systemPrompt ?? context.input.system ?? "";
const { text, object } = extractContent(context.result);
const output = stringifyContent(text, object, accumulatedText);
const result = detectSystemPromptLeak(output, systemPrompt, options);
if (!result.leaked) {
return {
tripwireTriggered: false,
info: { guardrailName: "system-prompt-leak-detector" }
};
}
return {
tripwireTriggered: true,
message: `System prompt leak detected (confidence: ${(result.confidence * 100).toFixed(1)}%): ${result.fragments.length} matching fragment(s)`,
severity: options.severity ?? "high",
metadata: {
confidence: result.confidence,
fragments: result.fragments,
sanitized: result.sanitized
},
suggestion: "The response reproduces the system prompt. Block it or replace it with metadata.sanitized.",
info: {
guardrailName: "system-prompt-leak-detector",
confidence: result.confidence,
fragmentCount: result.fragments.length
}
};
}
);
}
// src/guardrails/tool-parameters.ts
function matchesToolName(guardrailToolName, toolName) {
if (typeof guardrailToolName === "string") {
return guardrailToolName === toolName || guardrailToolName === "*";
}
if (guardrailToolName instanceof RegExp) {
return guardrailToolName.test(toolName);
}
if (Array.isArray(guardrailToolName)) {
return guardrailToolName.includes(toolName);
}
return false;
}
function withToolParameterGuardrails(tools, guardrails, options = {}) {
const { throwOnInvalid = true, onValidationFailed, requestContext } = options;
const wrappedTools = {};
for (const [toolName, tool] of Object.entries(tools)) {
const originalTool = tool;
const applicableGuardrails = guardrails.filter(
(g) => matchesToolName(g.toolName, toolName)
);
if (applicableGuardrails.length === 0) {
wrappedTools[toolName] = tool;
continue;
}
wrappedTools[toolName] = {
...originalTool,
execute: async (input, execOptions) => {
const context = {
toolName,
requestContext
};
let currentInput = input;
const failedResults = [];
for (const guardrail of applicableGuardrails) {
const result = await guardrail.validateInput(currentInput, context);
if (!result.valid) {
failedResults.push(result);
if (result.block) {
if (onValidationFailed) {
onValidationFailed(toolName, input, [result]);
}
if (throwOnInvalid) {
throw new ToolParameterValidationError(
toolName,
guardrail.name,
result.message || "Validation failed",
result.severity
);
}
return {
error: `Tool parameter validation failed: ${result.message}`,
blocked: true,
guardrail: guardrail.name
};
}
} else if (result.sanitizedInput !== void 0) {
currentInput = result.sanitizedInput;
}
}
if (failedResults.length > 0) {
if (onValidationFailed) {
onValidationFailed(toolName, input, failedResults);
}
if (throwOnInvalid) {
const messages = failedResults.map((r) => r.message).join("; ");
throw new ToolParameterValidationError(
toolName,
"multiple",
messages,
failedResults[0]?.severity
);
}
}
return originalTool.execute(currentInput, execOptions);
}
};
}
return wrappedTools;
}
var ToolParameterValidationError = class extends Error {
constructor(toolName, guardrailName, message, severity) {
super(
`Tool "${toolName}" parameter validation failed (${guardrailName}): ${message}`
);
this.toolName = toolName;
this.guardrailName = guardrailName;
this.severity = severity;
this.name = "ToolParameterValidationError";
}
toolName;
guardrailName;
severity;
};
function sqlInjectionGuardrail(options = {}) {
const defaultPatterns = [
/(\b(SELECT|INSERT|UPDATE|DELETE|DROP|UNION|ALTER|CREATE|TRUNCATE)\b.*\b(FROM|INTO|TABLE|DATABASE)\b)/i,
/(--)|(\/\*)|(\*\/)/,
/(\b(OR|AND)\b\s+\d+\s*=\s*\d+)/i,
/(;\s*(SELECT|INSERT|UPDATE|DELETE|DROP))/i,
/(\bEXEC\b|\bEXECUTE\b)/i
];
return {
name: "sql-injection-prevention",
description: "Prevents SQL injection attacks in tool parameters",
toolName: options.toolName || "*",
validateInput: (input) => {
const query = input.query || input.sql || "";
const patterns = options.patterns || defaultPatterns;
for (const pattern of patterns) {
if (pattern.test(query)) {
return {
valid: false,
block: true,
message: "Potential SQL injection detected",
severity: "critical",
metadata: { pattern: pattern.source }
};
}
}
return { valid: true };
}
};
}
function pathTraversalGuardrail(options = {}) {
const defaultBlockedPatterns = [
/\.\./,
/^\/etc\//,
/^\/proc\//,
/^\/sys\//,
/^~\//,
/\0/
];
return {
name: "path-traversal-prevention",
description: "Prevents path traversal attacks in file operations",
toolName: options.toolName || "*",
validateInput: (input) => {
const path = input.path || input.file || input.filename || "";
const blockedPatterns = options.blockedPatterns || defaultBlockedPatterns;
for (const pattern of blockedPatterns) {
if (pattern.test(path)) {
return {
valid: false,
block: true,
message: "Path traversal attempt detected",
severity: "critical",
metadata: { path, pattern: pattern.source }
};
}
}
if (options.allowedPaths && options.allowedPaths.length > 0) {
const isAllowed = options.allowedPaths.some(
(allowed) => path.startsWith(allowed)
);
if (!isAllowed) {
return {
valid: false,
block: true,
message: "Path not in allowed list",
severity: "high",
metadata: { path, allowedPaths: options.allowedPaths }
};
}
}
return { valid: true };
}
};
}
function parameterLengthGuardrail(options = {}) {
const maxLength = options.maxLength || 1e4;
return {
name: "parameter-length-limit",
description: "Enforces maximum length on tool parameters",
toolName: options.toolName || "*",
validateInput: (input) => {
const fields = options.fields || Object.keys(input);
for (const field of fields) {
const value = input[field];
if (typeof value === "string" && value.length > maxLength) {
return {
valid: false,
block: true,
message: `Parameter "${field}" exceeds maximum length (${value.length} > ${maxLength})`,
severity: "medium",
metadata: { field, length: value.length, maxLength }
};
}
}
return { valid: true };
}
};
}
function toolRBACGuardrail(options) {
const mode = options.mode || "any";
return {
name: "tool-rbac",
description: `Requires ${mode === "any" ? "any of" : "all of"} [${options.requiredPermissions.join(", ")}] permissions`,
toolName: options.toolName,
validateInput: (_input, context) => {
const userPermissions = context.requestContext?.permissions || [];
const hasPermission = mode === "any" ? options.requiredPermissions.some((p) => userPermissions.includes(p)) : options.requiredPermissions.every(
(p) => userPermissions.includes(p)
);
if (!hasPermission) {
return {
valid: false,
block: true,
message: `Insufficient permissions. Required: ${options.requiredPermissions.join(", ")}`,
severity: "high",
metadata: {
requiredPermissions: options.requiredPermissions,
userPermissions,
mode
}
};
}
return { valid: true };
}
};
}
export {
detectSystemPromptLeak,
systemPromptLeakDetector,
withToolParameterGuardrails,
ToolParameterValidationError,
sqlInjectionGuardrail,
pathTraversalGuardrail,
parameterLengthGuardrail,
toolRBACGuardrail
};