diff --git a/CHANGELOG.md b/CHANGELOG.md index bdf2f9c..276fe69 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,10 @@ # Changelog +## [Unreleased] + +- AI SDK 7 support, tested. `shieldLanguageModelMiddleware` works unchanged with AI SDK 7 (`ai@7`, `v4` language models) in `generateText` and `streamText`, and the test suite now runs it against AI SDK 4, 5, 6, and 7, including streaming, tool results, tool call arguments, `throwOnLeak`, and redaction. +- The legacy `shieldMiddleware().wrapParams()` and `wrapParamsAsync()` now also harden AI SDK 7's `instructions` option, which replaces the deprecated `system`. Before, a system prompt passed as `instructions` reached the model unhardened. + ## [2.0.0] - 2026-10-02 - Root `detect()` now returns a Promise and calls the hosted Shield API with a dashboard key. `createHostedDetector()` supports the four hosted model IDs, self-hosted moderation endpoints, cancellation, timeouts, and provider wrapper options. diff --git a/README.md b/README.md index b1d4f10..a7b6ceb 100644 --- a/README.md +++ b/README.md @@ -200,7 +200,7 @@ Every wrapper hardens the system prompt, runs its configured detector on user me | OpenAI | `shieldOpenAI` from `@zeroleaks/shield/openai` | `chat.completions.create`, `responses.create` | | Anthropic | `shieldAnthropic` from `@zeroleaks/shield/anthropic` | `messages.create` | | Groq | `shieldGroq` from `@zeroleaks/shield/groq` | `chat.completions.create` | -| Vercel AI SDK 4, 5, 6 | `shieldLanguageModelMiddleware` from `@zeroleaks/shield/ai-sdk` | `generateText`, `streamText` via `wrapLanguageModel` | +| Vercel AI SDK 4, 5, 6, 7 | `shieldLanguageModelMiddleware` from `@zeroleaks/shield/ai-sdk` | `generateText`, `streamText` via `wrapLanguageModel` | | Google Gen AI | `shieldGoogleGenAI` from `@zeroleaks/shield/google` | `models.generateContent`, `models.generateContentStream`, and chats | | Mistral | `shieldMistral` from `@zeroleaks/shield/mistral` | `chat.complete`, `chat.stream` | | LangChain.js | `shieldChatModel`, `ShieldCallbackHandler` from `@zeroleaks/shield/langchain` | `invoke`, `stream`, `batch`, and runnables derived from the model | diff --git a/bun.lock b/bun.lock index ad4f1ca..e585eb7 100644 --- a/bun.lock +++ b/bun.lock @@ -14,9 +14,10 @@ "@modelcontextprotocol/sdk": "^1.30.1", "@openai/agents": "^0.18.0", "@types/node": "^26.1.1", - "ai": "^6.0.292", + "ai": "^7.0.127", "ai-v4": "npm:ai@^4.3.19", "ai-v5": "npm:ai@^5.0.267", + "ai-v6": "npm:ai@^6.0.292", "groq-sdk": "^0.37.0", "openai": "^6.25.0", "tsup": "^8.5.1", @@ -51,7 +52,7 @@ }, }, "packages": { - "@ai-sdk/gateway": ["@ai-sdk/gateway@3.0.202", "", { "dependencies": { "@ai-sdk/provider": "3.0.17", "@ai-sdk/provider-utils": "4.0.54", "@vercel/oidc": "3.2.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-Ronhnf8CD4r1etZ37foVj+1klJfipLIcIYWH7/+WsotBrB64JdVzz2iC3PpzYYOcqKVcN1gvJPBYeJTCGfStRQ=="], + "@ai-sdk/gateway": ["@ai-sdk/gateway@4.0.103", "", { "dependencies": { "@ai-sdk/provider": "4.0.21", "@ai-sdk/provider-utils": "5.0.53", "@vercel/oidc": "3.2.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-nnGTUHzPRQ14jhv9CL4OQrzcRBIFTBMIPKGIi9Z9u/4GtwSRxa/nxUX6np2Iwp5phOzO0a0njj/j9Jolf5AEaA=="], "@ai-sdk/openai": ["@ai-sdk/openai@3.0.33", "", { "dependencies": { "@ai-sdk/provider": "3.0.8", "@ai-sdk/provider-utils": "4.0.15" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-O/8SVKAiwFHkGAUfBnrLb7L2IjbpP9ySWbmOktOfa0KtzutZkmKNrJ5CtB5dj+lwuENbOuZeRsnsZdOjar7hig=="], @@ -315,6 +316,8 @@ "@vitest/utils": ["@vitest/utils@4.0.18", "", { "dependencies": { "@vitest/pretty-format": "4.0.18", "tinyrainbow": "^3.0.3" } }, "sha512-msMRKLMVLWygpK3u2Hybgi4MNjcYJvwTb0Ru09+fOyCXIgT5raYP041DRRdiJiI3k/2U6SEbAETB3YtBrUkCFA=="], + "@workflow/serde": ["@workflow/serde@4.1.0", "", {}, "sha512-pav4F2BoirECWR7Nf1TKt+2eETcBj7jj4cBefQ8VXQCA6NPkaKeLfj/zMgi+3zYV5ZIBT4GuUiphsj0/b9hPQQ=="], + "abort-controller": ["abort-controller@3.0.0", "", { "dependencies": { "event-target-shim": "^5.0.0" } }, "sha512-h8lQ8tacZYnR3vNQTgibj+tODHI5/+l06Au2Pcriv/Gmet0eaj4TwWH41sO9wnHDiQsEj19q0drzdWdeAHtweg=="], "accepts": ["accepts@2.0.0", "", { "dependencies": { "mime-types": "^3.0.0", "negotiator": "^1.0.0" } }, "sha512-5cvg6CtKwfgdmVqY1WIiXKc3Q1bkRqGLi+2W/6ao+6Y7gu/RCwRuAhGEzh5B4KlszSuTLgZYuqFqo5bImjNKng=="], @@ -325,12 +328,14 @@ "agentkeepalive": ["agentkeepalive@4.6.0", "", { "dependencies": { "humanize-ms": "^1.2.1" } }, "sha512-kja8j7PjmncONqaTsB8fQ+wE2mSU2DJ9D4XKoJ5PFWIdRMa6SLSN1ff4mOr4jCbfRSsxR4keIiySJU0N9T5hIQ=="], - "ai": ["ai@6.0.292", "", { "dependencies": { "@ai-sdk/gateway": "3.0.202", "@ai-sdk/provider": "3.0.17", "@ai-sdk/provider-utils": "4.0.54", "@opentelemetry/api": "^1.9.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-bWHUQhqCoHUj+KGSmU6u5ikZY+3zbsLEBNGOxIYH8D+MLzWO/iuvFjeacE45ekfuBLX8Z1DqDnZsPL/to8pTiw=="], + "ai": ["ai@7.0.127", "", { "dependencies": { "@ai-sdk/gateway": "4.0.103", "@ai-sdk/provider": "4.0.21", "@ai-sdk/provider-utils": "5.0.53" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-JNsNPk4ZvRGEJuqdVo2lYFkm7uAuf3p/RwKFWn1Wh/+IK3U1RyFWw2GXoFnhzqNfxEuj7XJB4WpIUhl1sHToWg=="], "ai-v4": ["ai@4.3.19", "", { "dependencies": { "@ai-sdk/provider": "1.1.3", "@ai-sdk/provider-utils": "2.2.8", "@ai-sdk/react": "1.2.12", "@ai-sdk/ui-utils": "1.2.11", "@opentelemetry/api": "1.9.0", "jsondiffpatch": "0.6.0" }, "peerDependencies": { "react": "^18 || ^19 || ^19.0.0-rc", "zod": "^3.23.8" }, "optionalPeers": ["react"] }, "sha512-dIE2bfNpqHN3r6IINp9znguYdhIOheKW2LDigAMrgt/upT3B8eBGPSCblENvaZGoq+hxaN9fSMzjWpbqloP+7Q=="], "ai-v5": ["ai@5.0.267", "", { "dependencies": { "@ai-sdk/gateway": "2.0.159", "@ai-sdk/provider": "2.0.5", "@ai-sdk/provider-utils": "3.0.39", "@opentelemetry/api": "1.9.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-9oIQgz0v4elRrZ1mZJK83OJTH6SwV1aZD2nRUk6tgZ+DAnPDdLnaTtsKrbNB04Sb1aiEWUzkSaqfr/a9FDxBLg=="], + "ai-v6": ["ai@6.0.292", "", { "dependencies": { "@ai-sdk/gateway": "3.0.202", "@ai-sdk/provider": "3.0.17", "@ai-sdk/provider-utils": "4.0.54", "@opentelemetry/api": "^1.9.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-bWHUQhqCoHUj+KGSmU6u5ikZY+3zbsLEBNGOxIYH8D+MLzWO/iuvFjeacE45ekfuBLX8Z1DqDnZsPL/to8pTiw=="], + "ajv": ["ajv@8.20.0", "", { "dependencies": { "fast-deep-equal": "^3.1.3", "fast-uri": "^3.0.1", "json-schema-traverse": "^1.0.0", "require-from-string": "^2.0.2" } }, "sha512-Thbli+OlOj+iMPYFBVBfJ3OmCAnaSyNn4M1vz9T6Gka5Jt9ba/HIR56joy65tY6kx/FCF5VXNB819Y7/GUrBGA=="], "ajv-formats": ["ajv-formats@3.0.1", "", { "dependencies": { "ajv": "^8.0.0" } }, "sha512-8iUql50EUR+uUcdRQ3HDqa6EVyo3docL8g5WJ3FNcWmu62IbkGUue/pEyLBW8VGKKucTPgqeks4fIU1DA4yowQ=="], @@ -811,9 +816,9 @@ "zod-to-json-schema": ["zod-to-json-schema@3.25.1", "", { "peerDependencies": { "zod": "^3.25 || ^4" } }, "sha512-pM/SU9d3YAggzi6MtR4h7ruuQlqKtad8e9S0fmxcMi+ueAK5Korys/aWcV9LIIHTVbj01NdzxcnXSN+O74ZIVA=="], - "@ai-sdk/gateway/@ai-sdk/provider": ["@ai-sdk/provider@3.0.17", "", { "dependencies": { "json-schema": "^0.4.0" } }, "sha512-ijb+g+XxIAmv46hpt73ty83LhQHwwsqA1IwVX6p1MAREyX0zGM9eIrz1m3KlFxTHr8E4k/ao01w6Dlx3UBOihQ=="], + "@ai-sdk/gateway/@ai-sdk/provider": ["@ai-sdk/provider@4.0.21", "", { "dependencies": { "json-schema": "^0.4.0" } }, "sha512-UpbC9C1oht8dhfKPbVXSLRZS3DI8uk8n5v2uMjBPNIYGsC2kL045ywq7D8Hh9KgnH9rP5W/GxQdYtWScRouqlA=="], - "@ai-sdk/gateway/@ai-sdk/provider-utils": ["@ai-sdk/provider-utils@4.0.54", "", { "dependencies": { "@ai-sdk/provider": "3.0.17", "@standard-schema/spec": "^1.1.0", "eventsource-parser": "^3.0.8", "undici": "^6.28.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-piCQO5dnTqW7QSUbMsKg8oNI+u5VtgLnVFkgG6LfckrOyHcHFrIssBVJDg3/CXOdWXIFhVjIr540rzZAdom2TA=="], + "@ai-sdk/gateway/@ai-sdk/provider-utils": ["@ai-sdk/provider-utils@5.0.53", "", { "dependencies": { "@ai-sdk/provider": "4.0.21", "@standard-schema/spec": "^1.1.0", "@workflow/serde": "4.1.0", "eventsource-parser": "^3.0.8", "undici": "^7.29.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-VVe6UDd0y0/B4TfauuKKzcppqvhTmSluwwUVT4WG2/8h7gFbyFm7PM4UGM+ZXe3/RI6bv+Vl1cw/7aiBoQqdxg=="], "@ai-sdk/provider-utils/eventsource-parser": ["eventsource-parser@3.0.6", "", {}, "sha512-Vo1ab+QXPzZ4tCa8SwIHJFaSzy4R6SHf7BY79rFBDf0idraZWAkYrDjDj8uWaSm3S2TK+hJ7/t1CEmZ7jXw+pg=="], @@ -833,9 +838,9 @@ "@types/node-fetch/@types/node": ["@types/node@18.19.130", "", { "dependencies": { "undici-types": "~5.26.4" } }, "sha512-GRaXQx6jGfL8sKfaIDD6OupbIHBr9jv7Jnaml9tB7l4v068PAOXqfcujMMo5PhbIs6ggR1XODELqahT2R8v0fg=="], - "ai/@ai-sdk/provider": ["@ai-sdk/provider@3.0.17", "", { "dependencies": { "json-schema": "^0.4.0" } }, "sha512-ijb+g+XxIAmv46hpt73ty83LhQHwwsqA1IwVX6p1MAREyX0zGM9eIrz1m3KlFxTHr8E4k/ao01w6Dlx3UBOihQ=="], + "ai/@ai-sdk/provider": ["@ai-sdk/provider@4.0.21", "", { "dependencies": { "json-schema": "^0.4.0" } }, "sha512-UpbC9C1oht8dhfKPbVXSLRZS3DI8uk8n5v2uMjBPNIYGsC2kL045ywq7D8Hh9KgnH9rP5W/GxQdYtWScRouqlA=="], - "ai/@ai-sdk/provider-utils": ["@ai-sdk/provider-utils@4.0.54", "", { "dependencies": { "@ai-sdk/provider": "3.0.17", "@standard-schema/spec": "^1.1.0", "eventsource-parser": "^3.0.8", "undici": "^6.28.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-piCQO5dnTqW7QSUbMsKg8oNI+u5VtgLnVFkgG6LfckrOyHcHFrIssBVJDg3/CXOdWXIFhVjIr540rzZAdom2TA=="], + "ai/@ai-sdk/provider-utils": ["@ai-sdk/provider-utils@5.0.53", "", { "dependencies": { "@ai-sdk/provider": "4.0.21", "@standard-schema/spec": "^1.1.0", "@workflow/serde": "4.1.0", "eventsource-parser": "^3.0.8", "undici": "^7.29.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-VVe6UDd0y0/B4TfauuKKzcppqvhTmSluwwUVT4WG2/8h7gFbyFm7PM4UGM+ZXe3/RI6bv+Vl1cw/7aiBoQqdxg=="], "ai-v4/@ai-sdk/provider": ["@ai-sdk/provider@1.1.3", "", { "dependencies": { "json-schema": "^0.4.0" } }, "sha512-qZMxYJ0qqX/RfnuIaab+zp8UAeJn/ygXXAffR5I4N0n1IrvA6qBsjc8hXLmBiMV2zoXlifkacF7sEFnYnjBcqg=="], @@ -847,6 +852,12 @@ "ai-v5/@ai-sdk/provider-utils": ["@ai-sdk/provider-utils@3.0.39", "", { "dependencies": { "@ai-sdk/provider": "2.0.5", "@standard-schema/spec": "^1.0.0", "eventsource-parser": "^3.0.6", "undici": "^5.29.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-Cb1GE1UKcZxouWhV7weIQoJTgyz4WMswQfapPClN4QQP1pO01ZNSbG9EVN5C0gCsikA8vM0wrJZIT7HiZmH3wA=="], + "ai-v6/@ai-sdk/gateway": ["@ai-sdk/gateway@3.0.202", "", { "dependencies": { "@ai-sdk/provider": "3.0.17", "@ai-sdk/provider-utils": "4.0.54", "@vercel/oidc": "3.2.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-Ronhnf8CD4r1etZ37foVj+1klJfipLIcIYWH7/+WsotBrB64JdVzz2iC3PpzYYOcqKVcN1gvJPBYeJTCGfStRQ=="], + + "ai-v6/@ai-sdk/provider": ["@ai-sdk/provider@3.0.17", "", { "dependencies": { "json-schema": "^0.4.0" } }, "sha512-ijb+g+XxIAmv46hpt73ty83LhQHwwsqA1IwVX6p1MAREyX0zGM9eIrz1m3KlFxTHr8E4k/ao01w6Dlx3UBOihQ=="], + + "ai-v6/@ai-sdk/provider-utils": ["@ai-sdk/provider-utils@4.0.54", "", { "dependencies": { "@ai-sdk/provider": "3.0.17", "@standard-schema/spec": "^1.1.0", "eventsource-parser": "^3.0.8", "undici": "^6.28.0" }, "peerDependencies": { "zod": "^3.25.76 || ^4.1.8" } }, "sha512-piCQO5dnTqW7QSUbMsKg8oNI+u5VtgLnVFkgG6LfckrOyHcHFrIssBVJDg3/CXOdWXIFhVjIr540rzZAdom2TA=="], + "body-parser/content-type": ["content-type@2.1.0", "", {}, "sha512-mj7UPXE0jaqaOsukNZRUEfEi2AcL7C/vwmwcHV0O97eO1E1pxBZuyjlZrx5seTaNBg1U6+o35wpa35Qfcc+7ag=="], "form-data/mime-types": ["mime-types@2.1.35", "", { "dependencies": { "mime-db": "1.52.0" } }, "sha512-ZDY+bPm5zTTF+YpCrAU9nK0UgICYPT0QtT1NZWFv4s++TNkcgVaT0g6+4R2uI4MjQjzysHB1zxuWL50hzaeXiw=="], @@ -867,6 +878,8 @@ "vitest/tinyexec": ["tinyexec@1.0.2", "", {}, "sha512-W/KYk+NFhkmsYpuHq5JykngiOCnxeVL8v8dFnqxSD8qEEdRfXk1SDM6JzNqcERbcGYj9tMrDQBYV9cjgnunFIg=="], + "@ai-sdk/gateway/@ai-sdk/provider-utils/undici": ["undici@7.30.0", "", {}, "sha512-dkrQXeHSaoamnItlYbmzG0wFYrM0ZwDxCIg0A7aKjTyyhh9svRzCNFEzV+Vm05/yehjCzjDZ31KXfGEjYSztDQ=="], + "@ai-sdk/react/@ai-sdk/provider-utils/@ai-sdk/provider": ["@ai-sdk/provider@1.1.3", "", { "dependencies": { "json-schema": "^0.4.0" } }, "sha512-qZMxYJ0qqX/RfnuIaab+zp8UAeJn/ygXXAffR5I4N0n1IrvA6qBsjc8hXLmBiMV2zoXlifkacF7sEFnYnjBcqg=="], "@anthropic-ai/sdk/@types/node/undici-types": ["undici-types@5.26.5", "", {}, "sha512-JlCMO+ehdEIKqlFxk6IfVoAUVmgz7cU7zD/h9XZ0qzeosSHmUJVOzSQvvYSYWXkFXC+IfLKSIffhv0sVZup6pA=="], @@ -877,6 +890,8 @@ "ai-v5/@ai-sdk/provider-utils/undici": ["undici@5.29.0", "", { "dependencies": { "@fastify/busboy": "^2.0.0" } }, "sha512-raqeBD6NQK4SkWhQzeYKd1KmIG6dllBOTt55Rmkt4HtI9mwdWtJljnrXjAFUBLTSN67HWrOIZ3EPF4kjUw80Bg=="], + "ai/@ai-sdk/provider-utils/undici": ["undici@7.30.0", "", {}, "sha512-dkrQXeHSaoamnItlYbmzG0wFYrM0ZwDxCIg0A7aKjTyyhh9svRzCNFEzV+Vm05/yehjCzjDZ31KXfGEjYSztDQ=="], + "form-data/mime-types/mime-db": ["mime-db@1.52.0", "", {}, "sha512-sPU4uV7dYlvtWJxwwxHD0PuihVNiE7TyAbQ5SWxDCB9mUYvOgroQOwYQQOKPJ8CIbE+1ETVlOoK1UC2nU3gYvg=="], "groq-sdk/@types/node/undici-types": ["undici-types@5.26.5", "", {}, "sha512-JlCMO+ehdEIKqlFxk6IfVoAUVmgz7cU7zD/h9XZ0qzeosSHmUJVOzSQvvYSYWXkFXC+IfLKSIffhv0sVZup6pA=="], diff --git a/package.json b/package.json index 9c36ca2..5996cfe 100644 --- a/package.json +++ b/package.json @@ -173,9 +173,10 @@ "@modelcontextprotocol/sdk": "^1.30.1", "@openai/agents": "^0.18.0", "@types/node": "^26.1.1", - "ai": "^6.0.292", + "ai": "^7.0.127", "ai-v4": "npm:ai@^4.3.19", "ai-v5": "npm:ai@^5.0.267", + "ai-v6": "npm:ai@^6.0.292", "groq-sdk": "^0.37.0", "openai": "^6.25.0", "tsup": "^8.5.1", diff --git a/src/__tests__/ai-sdk.test.ts b/src/__tests__/ai-sdk.test.ts index 1da47fa..735d236 100644 --- a/src/__tests__/ai-sdk.test.ts +++ b/src/__tests__/ai-sdk.test.ts @@ -1,15 +1,11 @@ import { - generateText, - type ModelMessage, - type SystemModelMessage, - streamText, - wrapLanguageModel, + generateText as generateTextV7, + type ModelMessage as ModelMessageV7, + type SystemModelMessage as SystemModelMessageV7, + streamText as streamTextV7, + wrapLanguageModel as wrapLanguageModelV7, } from "ai"; -import { - convertArrayToReadableStream, - convertReadableStreamToArray, - MockLanguageModelV3, -} from "ai/test"; +import { MockLanguageModelV4 } from "ai/test"; import { generateText as generateTextV4, streamText as streamTextV4, @@ -23,6 +19,18 @@ import { wrapLanguageModel as wrapLanguageModelV5, } from "ai-v5"; import { MockLanguageModelV2 } from "ai-v5/test"; +import { + generateText, + type ModelMessage, + type SystemModelMessage, + streamText, + wrapLanguageModel, +} from "ai-v6"; +import { + convertArrayToReadableStream, + convertReadableStreamToArray, + MockLanguageModelV3, +} from "ai-v6/test"; import { describe, expect, it } from "vitest"; import { InjectionDetectedError, @@ -143,6 +151,83 @@ const V3_USAGE = { }; const V3_STOP = { unified: "stop" as const, raw: "stop" }; +/** AI SDK 7's `v4` models report usage and finish reasons as `v3` models do. */ +const V4_USAGE = V3_USAGE; +const V4_STOP = V3_STOP; + +function cleanV4(text = CLEAN) { + return new MockLanguageModelV4({ + doGenerate: { + content: [{ type: "text", text }], + finishReason: V4_STOP, + usage: V4_USAGE, + warnings: [], + }, + }); +} + +const aiSdk7: Harness = { + version: "7", + async generate(options, output, input = "Hi") { + const mock = new MockLanguageModelV4({ + doGenerate: { + content: [{ type: "text", text: output }], + finishReason: V4_STOP, + usage: V4_USAGE, + warnings: [], + response: { body: providerBody(output) }, + }, + }); + const result = await generateTextV7({ + model: wrapLanguageModelV7({ + model: mock, + middleware: shieldLanguageModelMiddleware(options), + }), + instructions: SYSTEM_PROMPT, + prompt: input, + // AI SDK 7 drops the response body unless asked to keep it. + include: { responseBody: true }, + }); + return { + text: result.text, + finishReason: result.finishReason, + system: systemOf(mock.doGenerateCalls[0].prompt), + body: result.response.body, + }; + }, + async stream(options, deltas, systemPrompt = SYSTEM_PROMPT) { + const errors: unknown[] = []; + const mock = new MockLanguageModelV4({ + doStream: { + stream: convertArrayToReadableStream([ + { type: "stream-start", warnings: [] }, + { type: "text-start", id: "t1" }, + ...deltas.map((delta) => ({ + type: "text-delta" as const, + id: "t1", + delta, + })), + { type: "text-end", id: "t1" }, + { type: "finish", finishReason: V4_STOP, usage: V4_USAGE }, + ]), + }, + }); + const result = streamTextV7({ + model: wrapLanguageModelV7({ + model: mock, + middleware: shieldLanguageModelMiddleware(options), + }), + instructions: systemPrompt, + prompt: "Hi", + onError: ({ error }) => { + errors.push(error); + }, + }); + const run = await readStream(result); + return { ...run, errors, system: systemOf(mock.doStreamCalls[0].prompt) }; + }, +}; + const aiSdk6: Harness = { version: "6", async generate(options, output, input = "Hi") { @@ -332,6 +417,7 @@ const aiSdk4: Harness = { }; describe.each([ + aiSdk7, aiSdk6, aiSdk5, aiSdk4, @@ -526,6 +612,152 @@ describe.each([ }); describe("shieldLanguageModelMiddleware stream parts", () => { + it("sanitizes each text block and keeps other parts in order on AI SDK 7", async () => { + const model = wrapLanguageModelV7({ + model: new MockLanguageModelV4({ + doStream: { + stream: convertArrayToReadableStream([ + { type: "stream-start", warnings: [] }, + { type: "reasoning-start", id: "r1" }, + { type: "reasoning-delta", id: "r1", delta: "Thinking." }, + { type: "reasoning-end", id: "r1" }, + { type: "text-start", id: "t1" }, + ...pieces(LEAKED, 9).map((delta) => ({ + type: "text-delta" as const, + id: "t1", + delta, + })), + { type: "text-end", id: "t1" }, + { + type: "tool-call", + toolCallId: "c1", + toolName: "lookup", + input: "{}", + }, + { type: "text-start", id: "t2" }, + { type: "text-delta", id: "t2", delta: CLEAN }, + { type: "text-end", id: "t2" }, + { type: "finish", finishReason: V4_STOP, usage: V4_USAGE }, + ]), + }, + }), + middleware: shieldLanguageModelMiddleware(), + }); + + const { stream } = await model.doStream({ + prompt: [ + { role: "system", content: SYSTEM_PROMPT }, + { role: "user", content: [{ type: "text", text: "Hi" }] }, + ], + }); + const parts = await convertReadableStreamToArray(stream); + const textOf = (id: string) => + parts + .map((part) => + part.type === "text-delta" && part.id === id ? part.delta : "" + ) + .join(""); + + expect(textOf("t1")).toBe(REDACTED_LEAK); + expect(textOf("t2")).toBe(CLEAN); + expect( + parts + .filter((part) => part.type !== "text-delta") + .map((part) => ("id" in part ? `${part.type}:${part.id}` : part.type)) + ).toEqual([ + "stream-start", + "reasoning-start:r1", + "reasoning-delta:r1", + "reasoning-end:r1", + "text-start:t1", + "text-end:t1", + "tool-call", + "text-start:t2", + "text-end:t2", + "finish", + ]); + }); + + it("keeps provider metadata carried on text deltas on AI SDK 7", async () => { + const providerMetadata = { google: { thoughtSignature: "sig123" } }; + const result = streamTextV7({ + model: wrapLanguageModelV7({ + model: new MockLanguageModelV4({ + doStream: { + stream: convertArrayToReadableStream([ + { type: "text-start", id: "t1" }, + { type: "text-delta", id: "t1", delta: CLEAN }, + { type: "text-delta", id: "t1", delta: "", providerMetadata }, + { type: "text-end", id: "t1" }, + { type: "finish", finishReason: V4_STOP, usage: V4_USAGE }, + ]), + }, + }), + middleware: shieldLanguageModelMiddleware(), + }), + instructions: SYSTEM_PROMPT, + prompt: "Hi", + }); + + expect(await result.content).toEqual([ + { type: "text", text: CLEAN, providerMetadata }, + ]); + }); + + it("sends the leak to a UI message stream as an error on AI SDK 7", async () => { + const result = streamTextV7({ + model: wrapLanguageModelV7({ + model: new MockLanguageModelV4({ + doStream: { + stream: convertArrayToReadableStream([ + { type: "text-start", id: "t1" }, + { type: "text-delta", id: "t1", delta: LEAKED }, + { type: "text-end", id: "t1" }, + { type: "finish", finishReason: V4_STOP, usage: V4_USAGE }, + ]), + }, + }), + middleware: shieldLanguageModelMiddleware({ throwOnLeak: true }), + }), + instructions: SYSTEM_PROMPT, + prompt: "Hi", + onError: () => undefined, + }); + + const body = await result.toUIMessageStreamResponse().text(); + + expect(body).toContain('"type":"error"'); + expect(body).toContain('"finishReason":"error"'); + expect(body).not.toContain("Never share account numbers"); + }); + + it("drops raw chunks on AI SDK 7, since they carry the unsanitized text", async () => { + const result = streamTextV7({ + model: wrapLanguageModelV7({ + model: new MockLanguageModelV4({ + doStream: { + stream: convertArrayToReadableStream([ + { type: "text-start", id: "t1" }, + { type: "raw", rawValue: providerBody(LEAKED) }, + { type: "text-delta", id: "t1", delta: LEAKED }, + { type: "text-end", id: "t1" }, + { type: "finish", finishReason: V4_STOP, usage: V4_USAGE }, + ]), + }, + }), + middleware: shieldLanguageModelMiddleware(), + }), + instructions: SYSTEM_PROMPT, + prompt: "Hi", + includeRawChunks: true, + }); + + const parts = await readAll(result.fullStream); + + expect(parts.map((part) => part.type)).not.toContain("raw"); + expect(JSON.stringify(parts)).not.toContain("Never share account numbers"); + }); + it("sanitizes each text block and keeps other parts in order", async () => { const model = wrapLanguageModel({ model: new MockLanguageModelV3({ @@ -811,6 +1043,117 @@ describe("shieldLanguageModelMiddleware without its own transformParams", () => }); }); +describe("shieldLanguageModelMiddleware wrapped by another middleware on AI SDK 7", () => { + it("sanitizes output when transformParams is wrapped", async () => { + const shield = shieldLanguageModelMiddleware(); + const result = await generateTextV7({ + model: wrapLanguageModelV7({ + model: cleanV4(LEAKED), + middleware: { + ...shield, + transformParams: async (options) => ({ + ...(await shield.transformParams(options)), + temperature: 0, + }), + }, + }), + instructions: SYSTEM_PROMPT, + prompt: "Hi", + }); + + expect(result.text).toBe(REDACTED_LEAK); + }); +}); + +describe("shieldMiddleware on AI SDK 7", () => { + it("hardens instructions for generateText and sanitizes the result", async () => { + const model = cleanV4(LEAKED); + const shield = shieldMiddleware({ systemPrompt: SYSTEM_PROMPT }); + + const result = await generateTextV7({ + model, + ...shield.wrapParams({ instructions: SYSTEM_PROMPT, prompt: "Hi" }), + }); + + expect(systemOf(model.doGenerateCalls[0].prompt)).toBe( + harden(SYSTEM_PROMPT) + ); + expect(shield.sanitizeOutput(result.text)).toBe(REDACTED_LEAK); + }); + + it("hardens the deprecated system option", async () => { + const model = cleanV4(CLEAN); + const shield = shieldMiddleware(); + + await generateTextV7({ + model, + ...shield.wrapParams({ system: SYSTEM_PROMPT, prompt: "Hi" }), + }); + + expect(systemOf(model.doGenerateCalls[0].prompt)).toBe( + harden(SYSTEM_PROMPT) + ); + }); + + it("hardens instructions for streamText", async () => { + const model = new MockLanguageModelV4({ + doStream: { + stream: convertArrayToReadableStream([ + { type: "text-start", id: "t1" }, + { type: "text-delta", id: "t1", delta: LEAKED }, + { type: "text-end", id: "t1" }, + { type: "finish", finishReason: V4_STOP, usage: V4_USAGE }, + ]), + }, + }); + const shield = shieldMiddleware({ systemPrompt: SYSTEM_PROMPT }); + const messages: ModelMessageV7[] = [{ role: "user", content: "Hi" }]; + + const result = streamTextV7({ + model, + ...shield.wrapParams({ instructions: SYSTEM_PROMPT, messages }), + }); + + expect(shield.sanitizeOutput(await result.text)).toBe(REDACTED_LEAK); + expect(systemOf(model.doStreamCalls[0].prompt)).toBe(harden(SYSTEM_PROMPT)); + }); + + it("hardens each system message in instructions and keeps their provider options", async () => { + const model = cleanV4(CLEAN); + const providerOptions = { + anthropic: { cacheControl: { type: "ephemeral" } }, + }; + const instructions: SystemModelMessageV7[] = [ + { role: "system", content: SYSTEM_PROMPT, providerOptions }, + { role: "system", content: "Answer in French." }, + ]; + const shield = shieldMiddleware(); + + await generateTextV7({ + model, + ...(await shield.wrapParamsAsync({ instructions, prompt: "Hi" })), + }); + + expect( + model.doGenerateCalls[0].prompt.filter( + (message) => message.role === "system" + ) + ).toEqual([ + { role: "system", content: harden(SYSTEM_PROMPT), providerOptions }, + { role: "system", content: harden("Answer in French.") }, + ]); + }); + + it("checks user messages passed as prompt", () => { + const shield = shieldMiddleware({ systemPrompt: SYSTEM_PROMPT }); + const prompt: ModelMessageV7[] = [{ role: "user", content: INJECTION }]; + + expect(() => + shield.wrapParams({ instructions: SYSTEM_PROMPT, prompt }) + ).toThrow(InjectionDetectedError); + }); +}); + describe("shieldMiddleware on AI SDK 6", () => { it("hardens params for generateText and sanitizes the result", async () => { const model = new MockLanguageModelV3({ @@ -924,7 +1267,7 @@ describe("shieldMiddleware on AI SDK 6", () => { }); }); -/** A user question, the model's tool call, and the tool's answer, as AI SDK 5 and 6 messages. */ +/** A user question, the model's tool call, and the tool's answer, as AI SDK 5, 6, and 7 messages. */ function toolTurn(output: unknown) { return [ { role: "user", content: "What's the weather in Paris?" }, @@ -977,6 +1320,49 @@ function cleanV3() { } describe("shieldLanguageModelMiddleware tool results", () => { + it.each( + TOOL_OUTPUTS + )("blocks an injection in a %s tool result on AI SDK 7", async (_, output) => { + const model = cleanV4(); + + const error = await rejection( + generateTextV7({ + model: wrapLanguageModelV7({ + model, + middleware: shieldLanguageModelMiddleware(), + }), + messages: toolTurn(output) as ModelMessageV7[], + }) + ); + + expect(error).toBeInstanceOf(InjectionDetectedError); + expect((error as InjectionDetectedError).source).toBe("tool"); + expect(model.doGenerateCalls).toHaveLength(0); + }); + + it("passes clean tool results and skips them with scanToolResults: false on AI SDK 7", async () => { + const clean = await generateTextV7({ + model: wrapLanguageModelV7({ + model: cleanV4(), + middleware: shieldLanguageModelMiddleware(), + }), + messages: toolTurn({ type: "text", value: "Sunny." }) as ModelMessageV7[], + }); + const skipped = await generateTextV7({ + model: wrapLanguageModelV7({ + model: cleanV4(), + middleware: shieldLanguageModelMiddleware({ scanToolResults: false }), + }), + messages: toolTurn({ + type: "text", + value: INJECTION, + }) as ModelMessageV7[], + }); + + expect(clean.text).toBe(CLEAN); + expect(skipped.text).toBe(CLEAN); + }); + it.each( TOOL_OUTPUTS )("blocks an injection in a %s tool result on AI SDK 6", async (_, output) => { @@ -1133,7 +1519,7 @@ describe("shieldLanguageModelMiddleware tool call arguments", () => { const args = JSON.stringify({ body: `Key ${AWS_KEY}` }); const safeArgs = JSON.stringify({ body: "Key [REDACTED]" }); - it("redacts a tool call's input on AI SDK 5 and 6", async () => { + it("redacts a tool call's input on AI SDK 5, 6, and 7", async () => { const middleware = shieldLanguageModelMiddleware(); const result = await middleware.wrapGenerate({ @@ -1192,6 +1578,84 @@ describe("shieldLanguageModelMiddleware tool call arguments", () => { ]); }); + it("redacts a tool call's input from generateText on AI SDK 7", async () => { + const model = new MockLanguageModelV4({ + doGenerate: { + content: [ + { type: "tool-call", toolCallId: "c1", toolName: "send", input: args }, + ], + finishReason: { unified: "tool-calls", raw: "tool_calls" }, + usage: V4_USAGE, + warnings: [], + response: { body: { raw: args } }, + }, + }); + + const result = await generateTextV7({ + model: wrapLanguageModelV7({ + model, + middleware: shieldLanguageModelMiddleware(), + }), + prompt: "Hi", + include: { responseBody: true }, + }); + + expect(result.toolCalls).toMatchObject([ + { toolCallId: "c1", toolName: "send", input: JSON.parse(safeArgs) }, + ]); + expect(result.response.body).toBeUndefined(); + }); + + it("redacts streamed tool input deltas and the tool call on AI SDK 7", async () => { + const model = wrapLanguageModelV7({ + model: new MockLanguageModelV4({ + doStream: { + stream: convertArrayToReadableStream([ + { type: "stream-start", warnings: [] }, + { type: "tool-input-start", id: "c1", toolName: "send" }, + ...pieces(args, 6).map((delta) => ({ + type: "tool-input-delta" as const, + id: "c1", + delta, + })), + { type: "tool-input-end", id: "c1" }, + { + type: "tool-call", + toolCallId: "c1", + toolName: "send", + input: args, + }, + { type: "finish", finishReason: V4_STOP, usage: V4_USAGE }, + ]), + }, + }), + middleware: shieldLanguageModelMiddleware(), + }); + + const { stream } = await model.doStream({ prompt }); + const parts = await convertReadableStreamToArray(stream); + + expect( + parts + .map((part) => (part.type === "tool-input-delta" ? part.delta : "")) + .join("") + ).toBe(safeArgs); + expect(parts.find((part) => part.type === "tool-call")).toMatchObject({ + input: safeArgs, + }); + expect( + parts + .filter((part) => part.type !== "tool-input-delta") + .map((part) => part.type) + ).toEqual([ + "stream-start", + "tool-input-start", + "tool-input-end", + "tool-call", + "finish", + ]); + }); + it("redacts streamed tool input deltas and the tool call on AI SDK 6", async () => { const model = wrapLanguageModel({ model: new MockLanguageModelV3({ diff --git a/src/providers/ai-sdk.ts b/src/providers/ai-sdk.ts index de85faf..b4a0f98 100644 --- a/src/providers/ai-sdk.ts +++ b/src/providers/ai-sdk.ts @@ -24,13 +24,19 @@ interface Message { role: string; content: string | MessagePart[]; } -/** AI SDK 6 also accepts system messages, alone or in an array, as `system`. */ +/** + * AI SDK 6 and later also accept system messages, alone or in an array, as + * `system`, and AI SDK 7 as `instructions`. + */ interface SystemMessage { role: "system"; content: string; } +type SystemParam = string | SystemMessage | Array; interface AISdkParams { - system?: string | SystemMessage | Array; + system?: SystemParam; + /** AI SDK 7's name for `system`, which it deprecates. */ + instructions?: SystemParam; /** AI SDK 5 and later also accept an array of messages here. */ prompt?: string | Array; messages?: Message[]; @@ -76,9 +82,9 @@ function hardenSystemMessage( * parts becomes a single hardened text part. */ function hardenSystem( - system: NonNullable, + system: SystemParam, options: HardenOptions -): AISdkParams["system"] { +): SystemParam { if (typeof system === "string") { return harden(system, options); } @@ -95,6 +101,23 @@ function hardenSystem( return text ? [{ type: "text", text: harden(text, options) }] : system; } +/** Hardens `system` and AI SDK 7's `instructions`, whichever are set. */ +function hardenParams

( + params: P, + options: HardenOptions | false +): P { + if (options === false) { + return { ...params }; + } + return { + ...params, + ...(params.system && { system: hardenSystem(params.system, options) }), + ...(params.instructions && { + instructions: hardenSystem(params.instructions, options), + }), + }; +} + /** Every message, whether in `messages` or in `prompt`. */ function paramMessages(params: AISdkParams): Message[] { const promptItems = Array.isArray(params.prompt) ? params.prompt : []; @@ -210,21 +233,12 @@ export function shieldMiddleware(options: ShieldAISdkOptions = {}) { ); } checkParams(params, shield.input); - - if (shield.harden === false || !params.system) { - return { ...params }; - } - return { - ...params, - system: hardenSystem(params.system, shield.harden), - }; + return hardenParams(params, shield.harden); }, async wrapParamsAsync

(params: P): Promise

{ await checkParamsAsync(params, shield.input); - return shield.harden === false || !params.system - ? { ...params } - : { ...params, system: hardenSystem(params.system, shield.harden) }; + return hardenParams(params, shield.harden); }, /** Redacts prompt leaks and output findings from `text`. */ @@ -291,7 +305,11 @@ interface LanguageModelStreamResult { stream: ReadableStream; } -/** Assignable to `LanguageModelMiddleware` from AI SDK 4, 5, and 6. */ +/** + * Assignable to `LanguageModelMiddleware` from AI SDK 4, 5, 6, and 7. AI SDK + * 7 accepts middleware of any specification version, and its `v4` call + * options, results, and stream parts have the shape this middleware reads. + */ export interface ShieldLanguageModelMiddleware { readonly specificationVersion: "v3"; transformParams:

(options: { @@ -649,7 +667,7 @@ function guardStreamParts( * AI SDK language model middleware. Pass it to `wrapLanguageModel` for * automatic hardening, injection detection on user input and tool results, * and output guarding in `generateText` and `streamText`, with no manual - * `sanitizeOutput` call. Works with AI SDK 4, 5, and 6. + * `sanitizeOutput` call. Works with AI SDK 4, 5, 6, and 7. * * @example * ```ts