UNPKG

ai-sdk-guardrails

Version:

Input and output guardrails middleware for Vercel AI SDK.

384 lines (380 loc) 13.2 kB
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 };