diff --git a/packages/apple-llm/ios/AppleLLM.mm b/packages/apple-llm/ios/AppleLLM.mm index 1b5d70d9..659cb59c 100644 --- a/packages/apple-llm/ios/AppleLLM.mm +++ b/packages/apple-llm/ios/AppleLLM.mm @@ -115,7 +115,8 @@ - (void)generateText:(nonnull NSArray *)messages @"topP": options.topP().has_value() ? @(options.topP().value()) : [NSNull null], @"topK": options.topK().has_value() ? @(options.topK().value()) : [NSNull null], @"schema": options.schema() ?: [NSNull null], - @"tools": options.tools() ?: [NSNull null] + @"tools": options.tools() ?: [NSNull null], + @"guardrails": options.guardrails() ?: [NSNull null] }; auto callToolBlock = ^(NSString *toolId, NSString *arguments, void (^completion)(id, NSError *)) { @@ -144,6 +145,7 @@ - (void)generateStream:(nonnull NSString *)streamId messages:(nonnull NSArray *) @"topK": options.topK().has_value() ? @(options.topK().value()) : [NSNull null], @"schema": options.schema() ?: [NSNull null], @"tools": options.tools() ?: [NSNull null], + @"guardrails": options.guardrails() ?: [NSNull null], }; auto callToolBlock = ^(NSString *toolId, NSString *arguments, void (^completion)(id, NSError *)) { diff --git a/packages/apple-llm/ios/AppleLLMImpl.swift b/packages/apple-llm/ios/AppleLLMImpl.swift index b1f0aed7..24327db4 100644 --- a/packages/apple-llm/ios/AppleLLMImpl.swift +++ b/packages/apple-llm/ios/AppleLLMImpl.swift @@ -88,7 +88,7 @@ public class AppleLLMImpl: NSObject { let (transcript, userPrompt) = try self.createTranscriptAndPrompt(from: messages, tools: tools) let session = LanguageModelSession.init( - model: SystemLanguageModel.default, + model: self.createModel(from: options), tools: tools, transcript: transcript ) @@ -155,7 +155,7 @@ public class AppleLLMImpl: NSObject { let (transcript, userPrompt) = try self.createTranscriptAndPrompt(from: messages, tools: tools) let session = LanguageModelSession.init( - model: SystemLanguageModel.default, + model: self.createModel(from: options), tools: tools, transcript: transcript ) @@ -384,6 +384,13 @@ public class AppleLLMImpl: NSObject { } @available(iOS 26, *) + private func createModel(from options: [String: Any]) -> SystemLanguageModel { + if options["guardrails"] as? String == "permissiveContentTransformations" { + return SystemLanguageModel(guardrails: .permissiveContentTransformations) + } + return SystemLanguageModel.default + } + private func createGenerationOptions(from options: [String: Any]) throws -> GenerationOptions { var temperature: Double? var maximumResponseTokens: Int? diff --git a/packages/apple-llm/src/NativeAppleLLM.ts b/packages/apple-llm/src/NativeAppleLLM.ts index 88a7456e..d38a3bcd 100644 --- a/packages/apple-llm/src/NativeAppleLLM.ts +++ b/packages/apple-llm/src/NativeAppleLLM.ts @@ -10,7 +10,11 @@ export interface AppleMessage { content: string } +export type AppleGuardrails = 'default' | 'permissiveContentTransformations' + export interface AppleGenerationOptions { + /** Guardrails mode for the underlying SystemLanguageModel. */ + guardrails?: AppleGuardrails temperature?: number maxTokens?: number topP?: number diff --git a/packages/apple-llm/src/ai-sdk.ts b/packages/apple-llm/src/ai-sdk.ts index d47e6539..0a4d2d34 100644 --- a/packages/apple-llm/src/ai-sdk.ts +++ b/packages/apple-llm/src/ai-sdk.ts @@ -23,7 +23,10 @@ import { import { createAppleLLMError, isAppleLLMErrorCode } from './errors' import NativeAppleEmbeddings from './NativeAppleEmbeddings' -import NativeAppleLLM, { type AppleMessage } from './NativeAppleLLM' +import NativeAppleLLM, { + type AppleGuardrails, + type AppleMessage, +} from './NativeAppleLLM' import NativeAppleSpeech from './NativeAppleSpeech' import NativeAppleTranscription from './NativeAppleTranscription' import NativeAppleUtils from './NativeAppleUtils' @@ -31,13 +34,24 @@ import NativeAppleUtils from './NativeAppleUtils' type Tool = LanguageModelV3FunctionTool | LanguageModelV3ProviderTool type ToolDefinitionSet = Record +export type { AppleGuardrails } from './NativeAppleLLM' + export function createAppleProvider({ availableTools, + guardrails, }: { availableTools?: ToolDefinitionSet + /** + * Guardrails mode for the on-device model. Use + * `'permissiveContentTransformations'` to lower guardrail sensitivity for + * apps that transform legitimate but sensitive content (for example health + * data). Defaults to the system default guardrails. + * @see https://developer.apple.com/documentation/foundationmodels/improving-the-safety-of-generative-model-output + */ + guardrails?: AppleGuardrails } = {}) { const createLanguageModel = () => { - return new AppleLLMChatLanguageModel(availableTools) + return new AppleLLMChatLanguageModel(availableTools, guardrails) } const provider = function () { return createLanguageModel() @@ -237,9 +251,14 @@ class AppleLLMChatLanguageModel implements LanguageModelV3 { readonly modelId = 'system-default' private tools: ToolDefinitionSet = {} + private guardrails?: AppleGuardrails - constructor(availableTools: ToolDefinitionSet = {}) { + constructor( + availableTools: ToolDefinitionSet = {}, + guardrails?: AppleGuardrails + ) { this.updateTools(availableTools) + this.guardrails = guardrails } async prepare(): Promise {} @@ -305,6 +324,7 @@ class AppleLLMChatLanguageModel implements LanguageModelV3 { try { const response = await NativeAppleLLM.generateText(messages, { + guardrails: this.guardrails, maxTokens: options.maxOutputTokens, temperature: options.temperature, topP: options.topP, @@ -365,6 +385,7 @@ class AppleLLMChatLanguageModel implements LanguageModelV3 { async doStream(options: LanguageModelV3CallOptions) { const messages = this.prepareMessages(options.prompt) const tools = this.prepareTools(options.tools) + const guardrails = this.guardrails if (typeof ReadableStream === 'undefined') { throw new Error( @@ -472,6 +493,7 @@ class AppleLLMChatLanguageModel implements LanguageModelV3 { listeners = [updateListener, completeListener, errorListener] NativeAppleLLM.generateStream(streamId, messages, { + guardrails, maxTokens: options.maxOutputTokens, temperature: options.temperature, topP: options.topP, diff --git a/website/src/docs/apple/generating.md b/website/src/docs/apple/generating.md index 07fda3c6..03494247 100644 --- a/website/src/docs/apple/generating.md +++ b/website/src/docs/apple/generating.md @@ -38,6 +38,24 @@ for await (const delta of textStream) { > [!NOTE] > Streaming objects is currently not supported. +## Guardrails + +Apple Foundation Models apply content guardrails to prompts and responses. Apps +that transform legitimate but sensitive content (for example health or medical +data) can hit `guardrailViolation` errors with the default mode. For those +use-cases, Apple provides a permissive guardrails mode that you can opt into +when creating the provider: + +```typescript +import { createAppleProvider } from '@react-native-ai/apple'; + +const apple = createAppleProvider({ + guardrails: 'permissiveContentTransformations' +}); +``` + +See [Improving the safety of generative model output](https://developer.apple.com/documentation/foundationmodels/improving-the-safety-of-generative-model-output#Use-permissive-guardrail-mode-for-sensitive-content) for when this mode is appropriate. + ## Structured Output Generate structured data that conforms to a specific schema: