diff --git a/CHANGELOG.md b/CHANGELOG.md index 2cc1f25..995bb7f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,27 @@ and Aria adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). ## [Unreleased] +## [0.13.0] - 2026-09-09 + +### Added + +- **Foundation Models Dynamic Profiles on iOS 27 and related platform + releases.** Applications can configure model selection, tools, instructions, + bounded history, transcript failure policy, response limits, and lifecycle + observability while retaining the existing iOS 26 session path. +- **Foundation Models multimodal prompts.** Text streaming and typed structured + generation now preserve validated in-memory JPEG and PNG image content and + explicitly request vision capability from system or injected models. +- **Typed Foundation Models failures.** Provider rejections expose stable + categories for unsupported capabilities, context limits, safety decisions, + and other failure policies without requiring consumers to inspect error text. + +### Changed + +- Foundation Models prompt, transcript, error, and session construction now use + shared conversion paths so text-only, multimodal, system-model, and custom + model execution behave consistently. + ## [0.12.0] - 2026-08-30 ### Added diff --git a/Sources/AgentKit/Runtime/AgentRuntime.swift b/Sources/AgentKit/Runtime/AgentRuntime.swift index 606b94a..5672fb4 100644 --- a/Sources/AgentKit/Runtime/AgentRuntime.swift +++ b/Sources/AgentKit/Runtime/AgentRuntime.swift @@ -278,6 +278,11 @@ import WorkflowKit return error.localizedDescription } switch agentError { + case let .providerRejected(failure): + if let underlying = failure.underlying { + return "Provider rejected (\(failure.kind.rawValue)): \(Self.condense(underlying.message))" + } + return "Provider rejected (\(failure.kind.rawValue)): \(failure.message)" case let .providerFailed(message, underlying): if let underlying { return "Provider failed: \(message) — \(Self.condense(underlying.message))" diff --git a/Sources/Aria/Aria.swift b/Sources/Aria/Aria.swift index 3abe056..4badbbe 100644 --- a/Sources/Aria/Aria.swift +++ b/Sources/Aria/Aria.swift @@ -21,7 +21,7 @@ public enum AriaInfo { /// /// Aria follows semantic versioning once it reaches `1.0.0`. Until then, /// breaking changes may occur on minor version bumps. - public static let version = "0.12.0" + public static let version = "0.13.0" } /// Internal logger used by core types. Backends are installed by the platform diff --git a/Sources/Aria/Foundation/AgentError.swift b/Sources/Aria/Foundation/AgentError.swift index 3f03475..c306016 100644 --- a/Sources/Aria/Foundation/AgentError.swift +++ b/Sources/Aria/Foundation/AgentError.swift @@ -7,6 +7,9 @@ import Foundation /// All recoverable errors travel as `AgentError` so callers can pattern-match /// instead of relying on string comparison or untyped `Error` values. public enum AgentError: Error, Sendable, Equatable { + /// The provider rejected a request with a stable, actionable category. + case providerRejected(ProviderFailure) + /// The provider failed for a reason the underlying SDK reported. The /// optional `underlying` carries the original error if available. case providerFailed(String, underlying: ErrorBox? = nil) @@ -38,6 +41,50 @@ public enum AgentError: Error, Sendable, Equatable { case configurationInvalid(String) } +// MARK: - ProviderFailureKind + +/// Provider-independent failure categories suitable for fallback policy and +/// observability. Providers translate their SDK-specific errors at the edge. +public enum ProviderFailureKind: String, Error, Sendable, Equatable { + case providerUnavailable + case assetsUnavailable + case sessionConflict + case transcriptMutation + case safetyRejected + case contextWindowExceeded + case unsupportedCapability + case unsupportedTranscript + case unsupportedGenerationGuide + case unsupportedLanguageOrLocale + case rateLimited + case timedOut + case invalidOutput + case unknown +} + +// MARK: - ProviderFailure + +/// A provider failure with a stable category and preserved SDK diagnostics. +public struct ProviderFailure: Sendable, Equatable { + // MARK: Lifecycle + + public init( + kind: ProviderFailureKind, + message: String, + underlying: ErrorBox? = nil + ) { + self.kind = kind + self.message = message + self.underlying = underlying + } + + // MARK: Public + + public let kind: ProviderFailureKind + public let message: String + public let underlying: ErrorBox? +} + // MARK: - ErrorBox /// A `Sendable` wrapper around an arbitrary `Error`. diff --git a/Sources/Aria/Foundation/Message.swift b/Sources/Aria/Foundation/Message.swift index 0b69ab7..f3de8b7 100644 --- a/Sources/Aria/Foundation/Message.swift +++ b/Sources/Aria/Foundation/Message.swift @@ -56,10 +56,9 @@ extension Message { /// Build a user message with one or more images attached. Each /// image is appended as a `ContentPart.image(...)` after the - /// text part. Vision-capable providers (FoundationModels with - /// vision, MLX VLM models) consume them; text-only providers - /// drop them silently via `textContent` (which only joins - /// `.text` parts). + /// text part. Vision-capable providers consume them; providers + /// without vision support reject or otherwise handle them according + /// to their capability policy. public static func user( _ text: String, images: [ImageContent], diff --git a/Sources/AriaApple/Providers/FoundationModelsDynamicProfileFactory.swift b/Sources/AriaApple/Providers/FoundationModelsDynamicProfileFactory.swift new file mode 100644 index 0000000..de19eee --- /dev/null +++ b/Sources/AriaApple/Providers/FoundationModelsDynamicProfileFactory.swift @@ -0,0 +1,210 @@ +#if canImport(FoundationModels) + import Aria + import FoundationModels + + public enum FoundationModelsTranscriptErrorPolicy: Sendable, Equatable { + case revert + case preserve + } + + public enum FoundationModelsProfileLifecycleEvent: Sendable, Equatable { + case activated + case deactivated + case prompt + case response + case reasoning + case toolCall + case toolOutput + } + + @available(iOS 26.0, macOS 26.0, *) + public struct FoundationModelsProfileConfiguration: Sendable { + // MARK: Lifecycle + + public init( + identifier: String, + maximumResponseTokens: Int? = nil, + historyLimit: Int? = nil, + transcriptErrorHandling: FoundationModelsTranscriptErrorPolicy = .revert, + lifecycleHandler: LifecycleHandler? = nil + ) { + self.identifier = identifier + self.maximumResponseTokens = maximumResponseTokens + self.historyLimit = historyLimit + self.transcriptErrorHandling = transcriptErrorHandling + self.lifecycleHandler = lifecycleHandler + } + + // MARK: Public + + public typealias LifecycleHandler = @Sendable ( + FoundationModelsProfileLifecycleEvent, + FoundationModelsProfileDescriptor + ) async -> Void + + public let identifier: String + public let maximumResponseTokens: Int? + public let historyLimit: Int? + public let transcriptErrorHandling: FoundationModelsTranscriptErrorPolicy + public let lifecycleHandler: LifecycleHandler? + } + + @available(iOS 26.0, macOS 26.0, *) + public struct FoundationModelsProfileDescriptor: Equatable, Sendable { + public let profileIdentifier: String + public let modelIdentifier: String + public let selectedToolNames: [String] + public let instructions: String? + public let maximumResponseTokens: Int? + public let historyLimit: Int? + public let transcriptErrorHandling: FoundationModelsTranscriptErrorPolicy + public let hasLifecycleHandler: Bool + } + + @available(iOS 26.0, macOS 26.0, *) + enum FoundationModelsDynamicProfileFactory { + // MARK: Internal + + static func validate(_ configuration: FoundationModelsProfileConfiguration) throws { + guard !configuration.identifier.isEmpty else { + throw AgentError.configurationInvalid( + "Foundation Models profile identifier must not be empty" + ) + } + if let maximumResponseTokens = configuration.maximumResponseTokens, + maximumResponseTokens <= 0 { + throw AgentError.configurationInvalid( + "Foundation Models maximum response tokens must be greater than zero" + ) + } + if let historyLimit = configuration.historyLimit, historyLimit < 0 { + throw AgentError.configurationInvalid( + "Foundation Models history limit must not be negative" + ) + } + } + + static func descriptor( + modelIdentifier: String, + tools: [any FoundationModels.Tool], + transcript: Transcript, + configuration: FoundationModelsProfileConfiguration + ) -> FoundationModelsProfileDescriptor { + FoundationModelsProfileDescriptor( + profileIdentifier: configuration.identifier, + modelIdentifier: modelIdentifier, + selectedToolNames: tools.map(\.name), + instructions: self.instructions(from: transcript), + maximumResponseTokens: configuration.maximumResponseTokens, + historyLimit: configuration.historyLimit, + transcriptErrorHandling: configuration.transcriptErrorHandling, + hasLifecycleHandler: configuration.lifecycleHandler != nil + ) + } + + static func history( + from transcript: Transcript, + limit: Int? + ) -> [Transcript.Entry] { + let entries = transcript.filter { entry in + if case .instructions = entry { + return false + } + return true + } + guard let limit else { + return entries + } + guard limit > 0 else { + return [] + } + var start = max(0, entries.count - limit) + if case .toolOutput = entries[start], + let toolCallsIndex = entries[.. String? { + let parts = transcript.compactMap { entry -> String? in + guard case let .instructions(instructions) = entry else { + return nil + } + return instructions.segments.compactMap { segment -> String? in + guard case let .text(text) = segment else { + return nil + } + return text.content + } + .joined() + } + .filter { !$0.isEmpty } + guard !parts.isEmpty else { + return nil + } + return parts.joined(separator: "\n\n") + } + } + + #if compiler(>=6.4) + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + @available(tvOS, unavailable) + extension FoundationModelsDynamicProfileFactory { + static func makeSession( + model: any LanguageModel, + modelIdentifier: String, + tools: [any FoundationModels.Tool], + transcript: Transcript, + configuration: FoundationModelsProfileConfiguration + ) throws -> LanguageModelSession { + try self.validate(configuration) + let descriptor = self.descriptor( + modelIdentifier: modelIdentifier, + tools: tools, + transcript: transcript, + configuration: configuration + ) + let instructionText = self.instructions(from: transcript) + let history = self.history( + from: transcript, + limit: configuration.historyLimit + ) + let policy: TranscriptErrorHandlingPolicy = + switch configuration.transcriptErrorHandling { + case .revert: .revertTranscript + case .preserve: .preserveTranscript + } + let handler = configuration.lifecycleHandler + let profile = LanguageModelSession.Profile { + if let instructionText { + Instructions(instructionText) + } + tools + } + .model(model) + .maximumResponseTokens(configuration.maximumResponseTokens) + .historyTransform { entries in + let transcript = Transcript(entries: entries) + return self.history(from: transcript, limit: configuration.historyLimit) + } + .transcriptErrorHandlingPolicy(policy) + .onPrompt { await handler?(.prompt, descriptor) } + .onResponse { await handler?(.response, descriptor) } + .onReasoning { await handler?(.reasoning, descriptor) } + .onToolCall { await handler?(.toolCall, descriptor) } + .onToolOutput { await handler?(.toolOutput, descriptor) } + .onActivate { await handler?(.activated, descriptor) } + .onDeactivate { await handler?(.deactivated, descriptor) } + return LanguageModelSession(profile: profile, history: history) + } + } + #endif +#endif diff --git a/Sources/AriaApple/Providers/FoundationModelsErrorMapper.swift b/Sources/AriaApple/Providers/FoundationModelsErrorMapper.swift new file mode 100644 index 0000000..b81a2d9 --- /dev/null +++ b/Sources/AriaApple/Providers/FoundationModelsErrorMapper.swift @@ -0,0 +1,150 @@ +#if canImport(FoundationModels) + import Aria + import Foundation + import FoundationModels + + @available(iOS 26.0, macOS 26.0, *) + enum FoundationModelsErrorMapper { + // MARK: Internal + + static func map(_ error: any Error) -> AgentError { + if error is CancellationError { + return .cancelled + } + if let agentError = error as? AgentError { + return agentError + } + + #if compiler(>=6.4) + if #available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *), + let mapped = mapCurrent(error) { + return mapped + } + #endif + + if let error = error as? LanguageModelSession.GenerationError { + return self.mapLegacy(error) + } + return self.rejected(.unknown, error) + } + + static func mapUnavailable( + _ reason: SystemLanguageModel.Availability.UnavailableReason + ) -> AgentError { + switch reason { + case .deviceNotEligible, .appleIntelligenceNotEnabled: + return self.rejected( + .providerUnavailable, + message: "Foundation Models is unavailable: \(String(describing: reason))" + ) + case .modelNotReady: + return self.rejected( + .assetsUnavailable, + message: "Foundation Models assets are not ready" + ) + @unknown default: + return self.rejected( + .providerUnavailable, + message: "Foundation Models availability is unknown" + ) + } + } + + // MARK: Private + + private static func mapLegacy( + _ error: LanguageModelSession.GenerationError + ) -> AgentError { + let kind: ProviderFailureKind = + switch error { + case .exceededContextWindowSize: + .contextWindowExceeded + case .assetsUnavailable: + .assetsUnavailable + case .guardrailViolation, .refusal: + .safetyRejected + case .unsupportedGuide: + .unsupportedGenerationGuide + case .unsupportedLanguageOrLocale: + .unsupportedLanguageOrLocale + case .decodingFailure: + .invalidOutput + case .rateLimited: + .rateLimited + case .concurrentRequests: + .sessionConflict + @unknown default: + .unknown + } + return self.rejected(kind, error) + } + + #if compiler(>=6.4) + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + private static func mapCurrent(_ error: any Error) -> AgentError? { + if let error = error as? LanguageModelError { + let kind: ProviderFailureKind = + switch error { + case .contextSizeExceeded: + .contextWindowExceeded + case .rateLimited: + .rateLimited + case .guardrailViolation, .refusal: + .safetyRejected + case .unsupportedCapability: + .unsupportedCapability + case .unsupportedTranscriptContent: + .unsupportedTranscript + case .unsupportedGenerationGuide: + .unsupportedGenerationGuide + case .unsupportedLanguageOrLocale: + .unsupportedLanguageOrLocale + case .timeout: + .timedOut + @unknown default: + .unknown + } + return self.rejected(kind, error) + } + if let error = error as? SystemLanguageModel.Error { + switch error { + case .assetsUnavailable: + return self.rejected(.assetsUnavailable, error) + @unknown default: + return self.rejected(.unknown, error) + } + } + if let error = error as? LanguageModelSession.Error { + switch error { + case .concurrentRequests: + return self.rejected(.sessionConflict, error) + case .transcriptMutationWhileResponding: + return self.rejected(.transcriptMutation, error) + @unknown default: + return self.rejected(.unknown, error) + } + } + return nil + } + #endif + + private static func rejected( + _ kind: ProviderFailureKind, + _ error: any Error + ) -> AgentError { + let underlying = ErrorBox(error) + return .providerRejected(.init( + kind: kind, + message: underlying.message, + underlying: underlying + )) + } + + private static func rejected( + _ kind: ProviderFailureKind, + message: String + ) -> AgentError { + .providerRejected(.init(kind: kind, message: message)) + } + } +#endif diff --git a/Sources/AriaApple/Providers/FoundationModelsImageBridge.swift b/Sources/AriaApple/Providers/FoundationModelsImageBridge.swift new file mode 100644 index 0000000..88b8db3 --- /dev/null +++ b/Sources/AriaApple/Providers/FoundationModelsImageBridge.swift @@ -0,0 +1,119 @@ +#if canImport(FoundationModels) && compiler(>=6.4) && (os(iOS) || os(macOS) || os(watchOS) || os(visionOS)) + import Aria + import CoreGraphics + import Foundation + import FoundationModels + import ImageIO + import UniformTypeIdentifiers + + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + enum FoundationModelsImageBridge { + // MARK: Internal + + enum PartKind: Equatable { + case text + case image + } + + enum ResolvedPart { + case text(String) + case image(CGImage) + + // MARK: Internal + + var kind: PartKind { + switch self { + case .text: .text + case .image: .image + } + } + + var text: String? { + guard case let .text(value) = self else { + return nil + } + return value + } + } + + static func resolve( + _ content: [ContentPart], + supportsVision: Bool + ) throws -> [ResolvedPart] { + try content.compactMap { part in + switch part { + case let .text(value): + return .text(value) + case let .image(image): + guard supportsVision else { + throw AgentError.providerRejected(.init( + kind: .unsupportedCapability, + message: "The selected Foundation Models model does not support image input" + )) + } + return try .image(self.decode(image)) + case .audio, .toolUse, .toolResult: + return nil + } + } + } + + static func prompt(from parts: [ResolvedPart]) -> Prompt { + let components = parts.map { part in + switch part { + case let .text(value): + Prompt(value) + case let .image(image): + Prompt(Attachment(image)) + } + } + return Prompt(components) + } + + static func transcriptSegments(from parts: [ResolvedPart]) -> [Transcript.Segment] { + parts.map { part in + switch part { + case let .text(value): + .text(.init(content: value)) + case let .image(image): + .attachment(.init(content: .image(.init(image)))) + } + } + } + + // MARK: Private + + private static func decode(_ image: ImageContent) throws -> CGImage { + guard case let .data(data, declaredMIMEType) = image.source else { + throw AgentError.configurationInvalid( + "Foundation Models image URLs and identifiers must be resolved to in-memory data by the host" + ) + } + + let normalizedMIMEType = declaredMIMEType + .trimmingCharacters(in: .whitespacesAndNewlines) + .lowercased() + let expectedType: UTType + switch normalizedMIMEType { + case "image/jpeg", "image/jpg": + expectedType = .jpeg + case "image/png": + expectedType = .png + default: + throw AgentError.configurationInvalid( + "Foundation Models supports in-memory JPEG and PNG image data; received \(declaredMIMEType)" + ) + } + + guard let source = CGImageSourceCreateWithData(data as CFData, nil), + let detectedIdentifier = CGImageSourceGetType(source), + detectedIdentifier as String == expectedType.identifier, + let decoded = CGImageSourceCreateImageAtIndex(source, 0, nil) else { + throw AgentError.configurationInvalid( + "Image data does not match its declared MIME type \(declaredMIMEType)" + ) + } + return decoded + } + } +#endif diff --git a/Sources/AriaApple/Providers/FoundationModelsProvider.swift b/Sources/AriaApple/Providers/FoundationModelsProvider.swift index 1da090c..40ac223 100644 --- a/Sources/AriaApple/Providers/FoundationModelsProvider.swift +++ b/Sources/AriaApple/Providers/FoundationModelsProvider.swift @@ -18,8 +18,10 @@ /// system prompts become `Instructions`, prior user turns become /// `prompt` entries, prior assistant text becomes `response` entries, /// prior tool calls become `toolCalls` entries, and prior tool - /// results become `toolOutput` entries. Only the *last* message in - /// the input array is sent as the new prompt to `streamResponse(to:)`. + /// results become `toolOutput` entries. On iOS 27, in-memory JPEG and + /// PNG content becomes an image attachment when the selected model + /// declares vision support. Only the *last* message in the input array + /// is sent as the new prompt to `streamResponse(to:)`. /// This avoids the transcript-style hallucination the model produces /// when given concatenated `User: …\nAssistant: …` text. @available(iOS 26.0, macOS 26.0, *) @@ -29,12 +31,14 @@ public init( defaultInstructions: String? = nil, capabilities: ProviderCapabilities = .foundationModelsDefault, - typedTools: [FoundationModelsToolFactory] = [] + typedTools: [FoundationModelsToolFactory] = [], + profileConfiguration: FoundationModelsProfileConfiguration? = nil ) { self.init( defaultInstructions: defaultInstructions, capabilities: capabilities, typedTools: typedTools, + profileConfiguration: profileConfiguration, sessionFactory: .systemDefault ) } @@ -46,12 +50,14 @@ model: some LanguageModel, defaultInstructions: String? = nil, capabilities: ProviderCapabilities, - typedTools: [FoundationModelsToolFactory] = [] + typedTools: [FoundationModelsToolFactory] = [], + profileConfiguration: FoundationModelsProfileConfiguration? = nil ) { self.init( defaultInstructions: defaultInstructions, capabilities: capabilities, typedTools: typedTools, + profileConfiguration: profileConfiguration, sessionFactory: .injected( model: model, declaredCapabilities: capabilities @@ -64,11 +70,13 @@ defaultInstructions: String? = nil, capabilities: ProviderCapabilities = .foundationModelsDefault, typedTools: [FoundationModelsToolFactory] = [], + profileConfiguration: FoundationModelsProfileConfiguration? = nil, sessionFactory: FoundationModelsSessionFactory ) { self.defaultInstructions = defaultInstructions self.capabilities = capabilities self.typedTools = typedTools + self.profileConfiguration = profileConfiguration self.sessionFactory = sessionFactory } @@ -114,12 +122,19 @@ // MARK: Internal + struct PreparedInput { + let prompt: Prompt + let transcript: Transcript + let requiresVision: Bool + } + static let maximumToolNameLength = 64 // Read by extensions in sibling files (e.g. // `FoundationModelsStructured.swift`). let defaultInstructions: String? let typedTools: [FoundationModelsToolFactory] + let profileConfiguration: FoundationModelsProfileConfiguration? let sessionFactory: FoundationModelsSessionFactory /// Pull the new-turn prompt out of the message list. Returns the @@ -146,6 +161,90 @@ return (last.textContent, Array(messages.dropLast())) } + static func extractPromptMessage( + from messages: [Message] + ) throws -> (prompt: Message, history: [Message]) { + guard let last = messages.last else { + throw AgentError.configurationInvalid( + "FoundationModelsProvider needs at least one message" + ) + } + let hasPromptContent = last.content.contains { part in + switch part { + case .text, .image: + true + case .audio, .toolUse, .toolResult: + false + } + } + guard hasPromptContent else { + throw AgentError.configurationInvalid( + "Last message must carry text or an image to seed the next response" + ) + } + return (last, Array(messages.dropLast())) + } + + static func containsImage(in content: [ContentPart]) -> Bool { + content.contains { part in + if case .image = part { + true + } else { + false + } + } + } + + static func prepareInput( + messages: [Message], + defaultInstructions: String?, + toolDefinitions: [Transcript.ToolDefinition], + supportsVision: Bool + ) throws -> PreparedInput { + let hasImages = messages.contains { self.containsImage(in: $0.content) } + guard hasImages else { + let extracted = try self.extractPrompt(from: messages) + return PreparedInput( + prompt: Prompt(extracted.prompt), + transcript: self.buildTranscript( + history: extracted.history, + defaultInstructions: defaultInstructions, + toolDefinitions: toolDefinitions + ), + requiresVision: false + ) + } + + #if compiler(>=6.4) + guard #available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) else { + throw AgentError.providerRejected(.init( + kind: .unsupportedCapability, + message: "Foundation Models image input requires iOS 27 or macOS 27" + )) + } + let extracted = try self.extractPromptMessage(from: messages) + let promptParts = try FoundationModelsImageBridge.resolve( + extracted.prompt.content, + supportsVision: supportsVision + ) + return try PreparedInput( + prompt: FoundationModelsImageBridge.prompt(from: promptParts), + transcript: self.buildMultimodalTranscript( + history: extracted.history, + defaultInstructions: defaultInstructions, + toolDefinitions: toolDefinitions, + supportsVision: supportsVision + ), + requiresVision: true + ) + #else + throw AgentError.providerRejected(.init( + kind: .unsupportedCapability, + message: "Foundation Models image input requires the iOS 27 SDK" + )) + #endif + } + /// Convert prior `Message` history into a `Transcript`. System /// messages collapse into a single `Instructions` entry (along /// with the configured default instructions and the bridge tool @@ -217,15 +316,12 @@ case .available: return case let .unavailable(reason): - throw AgentError.providerFailed( - "FoundationModels unavailable: \(String(describing: reason))", - underlying: nil - ) + throw FoundationModelsErrorMapper.mapUnavailable(reason) @unknown default: - throw AgentError.providerFailed( - "FoundationModels availability unknown", - underlying: nil - ) + throw AgentError.providerRejected(.init( + kind: .providerUnavailable, + message: "Foundation Models availability is unknown" + )) } } @@ -265,17 +361,8 @@ honourSelection: honourSelection, continuation: continuation ) - } catch is CancellationError { - continuation.finish(throwing: AgentError.cancelled) - } catch let error as AgentError { - continuation.finish(throwing: error) } catch { - continuation.finish( - throwing: AgentError.providerFailed( - "FoundationModels stream failed", - underlying: ErrorBox(error) - ) - ) + continuation.finish(throwing: FoundationModelsErrorMapper.map(error)) } } @@ -285,7 +372,6 @@ honourSelection: Bool, continuation: AsyncThrowingStream.Continuation ) async throws { - let (prompt, history) = try Self.extractPrompt(from: messages) // Build the typed FM tools. Each factory is invoked with a // closure that yields events into this stream's // continuation, so per-call `toolCallExecuted` events flow @@ -335,26 +421,32 @@ fmLog.error("tool not offered to FoundationModels: \(reason, privacy: .public)") } let toolDefinitions = registrableTools.map { Transcript.ToolDefinition(tool: $0) } - let transcript = Self.buildTranscript( - history: history, + let input = try Self.prepareInput( + messages: messages, defaultInstructions: self.defaultInstructions, - toolDefinitions: toolDefinitions + toolDefinitions: toolDefinitions, + supportsVision: self.capabilities.supportsVision ) var requirements: FoundationModelsSessionRequirements = [] if !registrableTools.isEmpty { requirements.insert(.toolCalling) } + if input.requiresVision { + requirements.insert(.vision) + } let session = try self.sessionFactory.makeSession( tools: registrableTools, - transcript: transcript, - requirements: requirements + transcript: input.transcript, + requirements: requirements, + modelIdentifier: self.capabilities.modelIdentifier, + profileConfiguration: self.profileConfiguration ) let messageId = UUID().uuidString continuation.yield(.messageStart(messageId: messageId)) var emittedCount = 0 - let stream = session.streamResponse(to: prompt) + let stream = session.streamResponse(to: input.prompt) for try await snapshot in stream { try Task.checkCancellation() // Extract the cumulative text synchronously inside the loop diff --git a/Sources/AriaApple/Providers/FoundationModelsSessionFactory.swift b/Sources/AriaApple/Providers/FoundationModelsSessionFactory.swift index 9a46773..e763b1a 100644 --- a/Sources/AriaApple/Providers/FoundationModelsSessionFactory.swift +++ b/Sources/AriaApple/Providers/FoundationModelsSessionFactory.swift @@ -31,25 +31,55 @@ typealias Builder = @Sendable ( [any FoundationModels.Tool], Transcript, - FoundationModelsSessionRequirements + FoundationModelsSessionRequirements, + String, + FoundationModelsProfileConfiguration? ) throws -> LanguageModelSession static let systemDefault = Self( validate: { _ in try FoundationModelsProvider.checkAvailability() }, - build: { tools, transcript, _ in - LanguageModelSession(tools: tools, transcript: transcript) + build: { tools, transcript, _, modelIdentifier, profileConfiguration in + #if compiler(>=6.4) + if #available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *), + let profileConfiguration { + return try FoundationModelsDynamicProfileFactory.makeSession( + model: SystemLanguageModel.default, + modelIdentifier: modelIdentifier, + tools: tools, + transcript: transcript, + configuration: profileConfiguration + ) + } + #endif + guard profileConfiguration == nil else { + throw AgentError.configurationInvalid( + "Foundation Models Dynamic Profiles require iOS 27 or macOS 27" + ) + } + return LanguageModelSession(tools: tools, transcript: transcript) } ) func makeSession( tools: [any FoundationModels.Tool], transcript: Transcript, - requirements: FoundationModelsSessionRequirements + requirements: FoundationModelsSessionRequirements, + modelIdentifier: String, + profileConfiguration: FoundationModelsProfileConfiguration? ) throws -> LanguageModelSession { try self.validate(requirements) - return try self.build(tools, transcript, requirements) + if let profileConfiguration { + try FoundationModelsDynamicProfileFactory.validate(profileConfiguration) + } + return try self.build( + tools, + transcript, + requirements, + modelIdentifier, + profileConfiguration + ) } // MARK: Private @@ -81,8 +111,17 @@ requested: requested ) }, - build: { tools, transcript, _ in - LanguageModelSession( + build: { tools, transcript, _, modelIdentifier, profileConfiguration in + if let profileConfiguration { + return try FoundationModelsDynamicProfileFactory.makeSession( + model: model, + modelIdentifier: modelIdentifier, + tools: tools, + transcript: transcript, + configuration: profileConfiguration + ) + } + return LanguageModelSession( model: model, tools: tools, transcript: transcript diff --git a/Sources/AriaApple/Providers/FoundationModelsStructured.swift b/Sources/AriaApple/Providers/FoundationModelsStructured.swift index 8e86c4b..5702db8 100644 --- a/Sources/AriaApple/Providers/FoundationModelsStructured.swift +++ b/Sources/AriaApple/Providers/FoundationModelsStructured.swift @@ -13,8 +13,9 @@ /// during generation, and concludes with `.finish(Content)`. /// /// The history portion of `messages` becomes the session - /// `Transcript`; the *last* element's text seeds the new - /// response (same convention as the text-streaming path). + /// `Transcript`; the *last* element's content seeds the new + /// response, including iOS 27 image attachments when supported + /// (same convention as the text-streaming path). public func streamStructured( messages: [Message], as type: Content.Type @@ -47,17 +48,8 @@ type: type, continuation: continuation ) - } catch is CancellationError { - continuation.finish(throwing: AgentError.cancelled) - } catch let error as AgentError { - continuation.finish(throwing: error) } catch { - continuation.finish( - throwing: AgentError.providerFailed( - "FoundationModels structured stream failed", - underlying: ErrorBox(error) - ) - ) + continuation.finish(throwing: FoundationModelsErrorMapper.map(error)) } } @@ -68,8 +60,6 @@ StructuredResponseEvent, any Error >.Continuation ) async throws where Content.PartiallyGenerated: Sendable { - let (prompt, history) = try Self.extractPrompt(from: messages) - // Forward each tool's `toolCallExecuted` ProviderEvent into // the structured stream so consumers see mid-response tool // activity. Other ProviderEvent cases would be @@ -83,22 +73,28 @@ } } let toolDefinitions = fmTools.map { Transcript.ToolDefinition(tool: $0) } - let transcript = Self.buildTranscript( - history: history, + let input = try Self.prepareInput( + messages: messages, defaultInstructions: self.defaultInstructions, - toolDefinitions: toolDefinitions + toolDefinitions: toolDefinitions, + supportsVision: self.capabilities.supportsVision ) var requirements: FoundationModelsSessionRequirements = [.guidedGeneration] if !fmTools.isEmpty { requirements.insert(.toolCalling) } + if input.requiresVision { + requirements.insert(.vision) + } let session = try self.sessionFactory.makeSession( tools: fmTools, - transcript: transcript, - requirements: requirements + transcript: input.transcript, + requirements: requirements, + modelIdentifier: self.capabilities.modelIdentifier, + profileConfiguration: self.profileConfiguration ) - let stream = session.streamResponse(to: prompt, generating: type) + let stream = session.streamResponse(to: input.prompt, generating: type) var lastRaw: GeneratedContent? for try await snapshot in stream { try Task.checkCancellation() diff --git a/Sources/AriaApple/Providers/FoundationModelsTranscript.swift b/Sources/AriaApple/Providers/FoundationModelsTranscript.swift index 446d3c1..43b2254 100644 --- a/Sources/AriaApple/Providers/FoundationModelsTranscript.swift +++ b/Sources/AriaApple/Providers/FoundationModelsTranscript.swift @@ -63,6 +63,85 @@ } } + #if compiler(>=6.4) + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + static func buildMultimodalTranscript( + history: [Message], + defaultInstructions: String?, + toolDefinitions: [Transcript.ToolDefinition], + supportsVision: Bool + ) throws -> Transcript { + if history.contains(where: { message in + message.role == .system && self.containsImage(in: message.content) + }) { + throw AgentError.configurationInvalid( + "Foundation Models does not accept image attachments in system instructions" + ) + } + + var entries: [Transcript.Entry] = [] + if let instructions = self.makeInstructions( + history: history, + defaultInstructions: defaultInstructions, + toolDefinitions: toolDefinitions + ) { + entries.append(.instructions(instructions)) + } + + let toolNames = self.toolNameMap(in: history) + for message in history where message.role != .system { + try entries.append(contentsOf: self.multimodalEntries( + for: message, + toolNames: toolNames, + supportsVision: supportsVision + )) + } + return Transcript(entries: entries) + } + + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + private static func multimodalEntries( + for message: Message, + toolNames: [String: String], + supportsVision: Bool + ) throws -> [Transcript.Entry] { + let parts = try FoundationModelsImageBridge.resolve( + message.content, + supportsVision: supportsVision + ) + let segments = FoundationModelsImageBridge.transcriptSegments(from: parts) + + switch message.role { + case .user: + guard !segments.isEmpty else { + return [] + } + return [.prompt(.init(segments: segments))] + case .assistant: + var entries: [Transcript.Entry] = [] + if !segments.isEmpty { + entries.append(.response(.init(assetIDs: [], segments: segments))) + } + let calls = message.toolCalls.compactMap(self.transcriptToolCall(from:)) + if !calls.isEmpty { + entries.append(.toolCalls(.init(calls))) + } + return entries + case .tool: + let id = message.toolCallId ?? UUID().uuidString + return [ + .toolOutput(.init( + id: id, + toolName: toolNames[id] ?? "unknown", + segments: segments + )), + ] + case .system: + return [] + } + } + #endif + // MARK: - Per-role entry builders private static func entriesForUser(_ message: Message) -> [Transcript.Entry] { @@ -105,7 +184,7 @@ return [.toolOutput(output)] } - private static func transcriptToolCall(from call: ToolCall) -> Transcript.ToolCall? { + static func transcriptToolCall(from call: ToolCall) -> Transcript.ToolCall? { guard let data = try? call.arguments.canonicalData(), let json = String(bytes: data, encoding: .utf8), let content = try? GeneratedContent(json: json) else { diff --git a/Tests/AriaAppleTests/Providers/FoundationModelsDynamicProfileFactoryTests.swift b/Tests/AriaAppleTests/Providers/FoundationModelsDynamicProfileFactoryTests.swift new file mode 100644 index 0000000..dd1d8e8 --- /dev/null +++ b/Tests/AriaAppleTests/Providers/FoundationModelsDynamicProfileFactoryTests.swift @@ -0,0 +1,205 @@ +#if canImport(FoundationModels) && (os(iOS) || os(macOS) || os(watchOS) || os(tvOS) || os(visionOS)) + import Aria + @testable import AriaApple + import Foundation + import FoundationModels + import XCTest + + @available(iOS 26.0, macOS 26.0, *) + final class FoundationModelsDynamicProfileFactoryTests: XCTestCase { + override func setUpWithError() throws { + guard #available(iOS 26.0, macOS 26.0, *) else { + throw XCTSkip("Requires iOS 26 / macOS 26 runtime") + } + } + + func testDescriptorCapturesEffectiveProfileInputs() { + let transcript = FoundationModelsProvider.buildTranscript( + history: [ + .system("Classify the meal."), + .user("A bowl of lentil soup."), + ], + defaultInstructions: "Return one category.", + toolDefinitions: [] + ) + let configuration = FoundationModelsProfileConfiguration( + identifier: "food-classifier-v1", + maximumResponseTokens: 128, + historyLimit: 1, + transcriptErrorHandling: .preserve, + lifecycleHandler: { _, _ in } + ) + + let descriptor = FoundationModelsDynamicProfileFactory.descriptor( + modelIdentifier: "apple.foundationmodels.classifier", + tools: [SessionFactoryTestTool()], + transcript: transcript, + configuration: configuration + ) + + XCTAssertEqual(descriptor.profileIdentifier, "food-classifier-v1") + XCTAssertEqual(descriptor.modelIdentifier, "apple.foundationmodels.classifier") + XCTAssertEqual(descriptor.selectedToolNames, ["session_factory_test"]) + XCTAssertEqual( + descriptor.instructions, + "Return one category.\n\nClassify the meal." + ) + XCTAssertEqual(descriptor.maximumResponseTokens, 128) + XCTAssertEqual(descriptor.historyLimit, 1) + XCTAssertEqual(descriptor.transcriptErrorHandling, .preserve) + XCTAssertTrue(descriptor.hasLifecycleHandler) + } + + func testHistoryDropsInstructionsAndKeepsConfiguredSuffix() { + let transcript = FoundationModelsProvider.buildTranscript( + history: [ + .system("Use the requested format."), + .user("first"), + .assistant("second"), + .user("third"), + ], + defaultInstructions: nil, + toolDefinitions: [] + ) + + let history = FoundationModelsDynamicProfileFactory.history( + from: transcript, + limit: 2 + ) + + XCTAssertEqual(history.count, 2) + XCTAssertEqual(entryKind(history[0]), "response") + XCTAssertEqual(entryKind(history[1]), "prompt") + } + + func testNilHistoryLimitKeepsAllNonInstructionEntries() { + let transcript = FoundationModelsProvider.buildTranscript( + history: [ + .system("Be concise."), + .user("first"), + .assistant("second"), + ], + defaultInstructions: nil, + toolDefinitions: [] + ) + + let history = FoundationModelsDynamicProfileFactory.history( + from: transcript, + limit: nil + ) + + XCTAssertEqual(history.count, 2) + XCTAssertEqual(entryKind(history[0]), "prompt") + XCTAssertEqual(entryKind(history[1]), "response") + } + + func testHistoryLimitDoesNotOrphanToolOutput() { + let call = ToolCall( + id: "weather-1", + name: "weather", + arguments: .object(["city": .string("Cupertino")]) + ) + let transcript = FoundationModelsProvider.buildTranscript( + history: [ + .user("Check the weather."), + .assistant("", toolCalls: [call]), + .tool(callId: call.id, text: "Sunny"), + .assistant("It is sunny."), + ], + defaultInstructions: nil, + toolDefinitions: [] + ) + + let history = FoundationModelsDynamicProfileFactory.history( + from: transcript, + limit: 2 + ) + + XCTAssertEqual(history.count, 3) + XCTAssertEqual(entryKind(history[0]), "toolCalls") + XCTAssertEqual(entryKind(history[1]), "toolOutput") + XCTAssertEqual(entryKind(history[2]), "response") + } + + func testInvalidProfileValuesAreRejectedBeforeSessionConstruction() { + let configurations = [ + FoundationModelsProfileConfiguration(identifier: ""), + FoundationModelsProfileConfiguration( + identifier: "invalid-response-limit", + maximumResponseTokens: 0 + ), + FoundationModelsProfileConfiguration( + identifier: "invalid-history-limit", + historyLimit: -1 + ), + ] + + for configuration in configurations { + XCTAssertThrowsError( + try FoundationModelsDynamicProfileFactory.validate(configuration) + ) { error in + guard case AgentError.configurationInvalid = error else { + return XCTFail("Expected configurationInvalid, got \(error)") + } + } + } + } + + #if compiler(>=6.4) + #if !os(tvOS) + func testDynamicProfileConstructsARealSession() throws { + guard ProcessInfo.processInfo.environment["ARIA_RUN_EVALS"] == "1" else { + throw XCTSkip("Runs a real model; set ARIA_RUN_EVALS=1") + } + guard #available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) else { + throw XCTSkip("Requires iOS 27 / macOS 27 runtime") + } + guard SystemLanguageModel.default.availability == .available else { + throw XCTSkip("Requires available Foundation Models assets") + } + try self.assertDynamicProfileConstructsARealSession() + } + + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + private func assertDynamicProfileConstructsARealSession() throws { + let transcript = FoundationModelsProvider.buildTranscript( + history: [.user("previous prompt"), .assistant("previous response")], + defaultInstructions: "Be concise.", + toolDefinitions: [] + ) + + let session = try FoundationModelsDynamicProfileFactory.makeSession( + model: SystemLanguageModel.default, + modelIdentifier: "apple.foundationmodels.default", + tools: [], + transcript: transcript, + configuration: .init( + identifier: "dynamic-profile-construction", + maximumResponseTokens: 64, + historyLimit: 1, + transcriptErrorHandling: .revert + ) + ) + + XCTAssertEqual(session.transcript.count, 2) + XCTAssertEqual(entryKind(session.transcript[0]), "instructions") + XCTAssertEqual(entryKind(session.transcript[1]), "response") + } + #endif + #endif + + private func entryKind(_ entry: Transcript.Entry) -> String { + switch entry { + case .instructions: "instructions" + case .prompt: "prompt" + case .toolCalls: "toolCalls" + case .toolOutput: "toolOutput" + case .response: "response" + #if compiler(>=6.4) + case .reasoning: "reasoning" + #endif + @unknown default: "unknown" + } + } + } +#endif diff --git a/Tests/AriaAppleTests/Providers/FoundationModelsErrorMapperTests.swift b/Tests/AriaAppleTests/Providers/FoundationModelsErrorMapperTests.swift new file mode 100644 index 0000000..fe09ada --- /dev/null +++ b/Tests/AriaAppleTests/Providers/FoundationModelsErrorMapperTests.swift @@ -0,0 +1,171 @@ +#if canImport(FoundationModels) && (os(iOS) || os(macOS) || os(watchOS) || os(tvOS) || os(visionOS)) + + import Foundation + import FoundationModels + import XCTest + @testable import Aria + @testable import AriaApple + + @available(iOS 26.0, macOS 26.0, *) + final class FoundationModelsErrorMapperTests: XCTestCase { + override func setUpWithError() throws { + guard #available(iOS 26.0, macOS 26.0, *) else { + throw XCTSkip("Requires iOS 26 / macOS 26 runtime") + } + } + + func testCancellationStaysCancellation() { + XCTAssertEqual( + FoundationModelsErrorMapper.map(CancellationError()), + .cancelled + ) + } + + func testExistingAgentErrorPassesThroughUnchanged() { + let expected = AgentError.configurationInvalid("bad configuration") + XCTAssertEqual(FoundationModelsErrorMapper.map(expected), expected) + } + + func testUnknownErrorPreservesTypeAndMessage() { + let mapped = FoundationModelsErrorMapper.map(TestError.failure) + guard case let .providerRejected(failure) = mapped else { + return XCTFail("Expected typed provider rejection, got \(mapped)") + } + + XCTAssertEqual(failure.kind, .unknown) + XCTAssertEqual(failure.underlying?.typeName, "TestError") + XCTAssertTrue(failure.underlying?.message.contains("failure") == true) + } + + func testUnavailableReasonsMapToActionableKinds() { + XCTAssertEqual( + FoundationModelsErrorMapper.mapUnavailable(.deviceNotEligible).failureKind, + .providerUnavailable + ) + XCTAssertEqual( + FoundationModelsErrorMapper.mapUnavailable(.appleIntelligenceNotEnabled).failureKind, + .providerUnavailable + ) + XCTAssertEqual( + FoundationModelsErrorMapper.mapUnavailable(.modelNotReady).failureKind, + .assetsUnavailable + ) + } + + #if compiler(>=6.4) + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + func testLanguageModelErrorsMapToStableKinds() throws { + guard #available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) else { + throw XCTSkip("Requires iOS 27 / macOS 27 runtime") + } + let cases: [(any Error, ProviderFailureKind)] = [ + ( + LanguageModelError.contextSizeExceeded(.init( + contextSize: 4096, + tokenCount: 4097, + debugDescription: "too many tokens" + )), + .contextWindowExceeded + ), + ( + LanguageModelError.rateLimited(.init( + resetDate: nil, + debugDescription: "slow down" + )), + .rateLimited + ), + ( + LanguageModelError.guardrailViolation(.init( + debugDescription: "guardrail" + )), + .safetyRejected + ), + ( + LanguageModelError.refusal(.init( + explanation: "cannot comply", + debugDescription: "refused" + )), + .safetyRejected + ), + ( + LanguageModelError.unsupportedCapability(.init( + capability: .vision, + debugDescription: "no images" + )), + .unsupportedCapability + ), + ( + LanguageModelError.unsupportedTranscriptContent(.init( + unsupportedContent: [], + debugDescription: "bad transcript" + )), + .unsupportedTranscript + ), + ( + LanguageModelError.unsupportedGenerationGuide(.init( + schemaName: "Food", + debugDescription: "bad guide" + )), + .unsupportedGenerationGuide + ), + ( + LanguageModelError.unsupportedLanguageOrLocale(.init( + languageCode: .english, + debugDescription: "unsupported language" + )), + .unsupportedLanguageOrLocale + ), + ( + LanguageModelError.timeout(.init(debugDescription: "timed out")), + .timedOut + ), + ] + + for (error, expectedKind) in cases { + XCTAssertEqual( + FoundationModelsErrorMapper.map(error).failureKind, + expectedKind + ) + } + } + + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + func testSystemAndSessionErrorsMapToStableKinds() throws { + guard #available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) else { + throw XCTSkip("Requires iOS 27 / macOS 27 runtime") + } + let cases: [(any Error, ProviderFailureKind)] = [ + ( + SystemLanguageModel.Error.assetsUnavailable(.init( + debugDescription: "assets missing" + )), + .assetsUnavailable + ), + (LanguageModelSession.Error.concurrentRequests, .sessionConflict), + (LanguageModelSession.Error.transcriptMutationWhileResponding, .transcriptMutation), + ] + + for (error, expectedKind) in cases { + XCTAssertEqual( + FoundationModelsErrorMapper.map(error).failureKind, + expectedKind + ) + } + } + #endif + } + + private extension AgentError { + var failureKind: ProviderFailureKind? { + guard case let .providerRejected(failure) = self else { + return nil + } + return failure.kind + } + } + + private enum TestError: Error { + case failure + } + +#endif diff --git a/Tests/AriaAppleTests/Providers/FoundationModelsImageBridgeTests.swift b/Tests/AriaAppleTests/Providers/FoundationModelsImageBridgeTests.swift new file mode 100644 index 0000000..59e4fb9 --- /dev/null +++ b/Tests/AriaAppleTests/Providers/FoundationModelsImageBridgeTests.swift @@ -0,0 +1,186 @@ +#if canImport(FoundationModels) && compiler(>=6.4) && (os(iOS) || os(macOS) || os(watchOS) || os(visionOS)) + import Aria + @testable import AriaApple + import CoreGraphics + import Foundation + import FoundationModels + import ImageIO + import UniformTypeIdentifiers + import XCTest + + @available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) + final class FoundationModelsImageBridgeTests: XCTestCase { + override func setUpWithError() throws { + guard #available(iOS 27.0, macOS 27.0, visionOS 27.0, watchOS 27.0, *) else { + throw XCTSkip("Requires iOS 27 / macOS 27 runtime") + } + } + + func testJPEGAndPNGDataResolveAsImages() throws { + let parts = try FoundationModelsImageBridge.resolve( + [ + .image(.init(source: .data(try imageData(type: .jpeg), mimeType: "image/jpeg"))), + .image(.init(source: .data(try imageData(type: .png), mimeType: "image/png"))), + ], + supportsVision: true + ) + + XCTAssertEqual(parts.map(\.kind), [.image, .image]) + } + + func testMixedContentPreservesTextAndImageOrder() throws { + let parts = try FoundationModelsImageBridge.resolve( + [ + .text("before"), + .image(.init(source: .data(try imageData(type: .png), mimeType: "image/png"))), + .text("after"), + ], + supportsVision: true + ) + + XCTAssertEqual(parts.map(\.kind), [.text, .image, .text]) + XCTAssertEqual(parts.compactMap(\.text), ["before", "after"]) + } + + func testInvalidMIMETypeIsRejected() throws { + XCTAssertThrowsError( + try FoundationModelsImageBridge.resolve( + [.image(.init(source: .data(try imageData(type: .png), mimeType: "image/gif")))], + supportsVision: true + ) + ) { error in + assertConfigurationInvalid(error) + } + } + + func testDeclaredMIMEMustMatchImageData() throws { + XCTAssertThrowsError( + try FoundationModelsImageBridge.resolve( + [.image(.init(source: .data(try imageData(type: .jpeg), mimeType: "image/png")))], + supportsVision: true + ) + ) { error in + assertConfigurationInvalid(error) + } + } + + func testURLAndIdentifierRequireHostResolution() throws { + let sources: [ImageContent.Source] = [ + .url(try XCTUnwrap(URL(string: "https://example.com/meal.jpg"))), + .identifier("photos-local-identifier"), + ] + + for source in sources { + XCTAssertThrowsError( + try FoundationModelsImageBridge.resolve( + [.image(.init(source: source))], + supportsVision: true + ) + ) { error in + assertConfigurationInvalid(error) + } + } + } + + func testTextOnlyModelRejectsImageWithTypedFailure() throws { + XCTAssertThrowsError( + try FoundationModelsImageBridge.resolve( + [.image(.init(source: .data(try imageData(type: .png), mimeType: "image/png")))], + supportsVision: false + ) + ) { error in + guard case let AgentError.providerRejected(failure) = error else { + return XCTFail("Expected providerRejected, got \(error)") + } + XCTAssertEqual(failure.kind, .unsupportedCapability) + } + } + + func testPreparedInputAcceptsAnImageOnlyPromptAndRequiresVision() throws { + let input = try FoundationModelsProvider.prepareInput( + messages: [Message( + role: .user, + content: [.image(.init(source: .data( + try imageData(type: .png), + mimeType: "image/png" + )))] + )], + defaultInstructions: nil, + toolDefinitions: [], + supportsVision: true + ) + + XCTAssertTrue(input.requiresVision) + XCTAssertTrue(input.transcript.isEmpty) + } + + func testHistoricalImageBecomesAnAttachmentSegmentInOrder() throws { + let input = try FoundationModelsProvider.prepareInput( + messages: [ + Message( + role: .user, + content: [ + .text("before"), + .image(.init(source: .data( + try imageData(type: .jpeg), + mimeType: "image/jpeg" + ))), + .text("after"), + ] + ), + .user("continue"), + ], + defaultInstructions: nil, + toolDefinitions: [], + supportsVision: true + ) + + let entry = try XCTUnwrap(input.transcript.first) + guard case let .prompt(prompt) = entry else { + return XCTFail("Expected historical prompt") + } + XCTAssertEqual(prompt.segments.count, 3) + guard case .text = prompt.segments[0], + case .attachment = prompt.segments[1], + case .text = prompt.segments[2] else { + return XCTFail("Expected text, image, text transcript ordering") + } + } + + private func assertConfigurationInvalid( + _ error: any Error, + file: StaticString = #filePath, + line: UInt = #line + ) { + guard case AgentError.configurationInvalid = error else { + return XCTFail("Expected configurationInvalid, got \(error)", file: file, line: line) + } + } + + private func imageData(type: UTType) throws -> Data { + let colorSpace = try XCTUnwrap(CGColorSpace(name: CGColorSpace.sRGB)) + let context = try XCTUnwrap(CGContext( + data: nil, + width: 1, + height: 1, + bitsPerComponent: 8, + bytesPerRow: 4, + space: colorSpace, + bitmapInfo: CGImageAlphaInfo.premultipliedLast.rawValue + )) + context.setFillColor(CGColor(red: 1, green: 0, blue: 0, alpha: 1)) + context.fill(CGRect(x: 0, y: 0, width: 1, height: 1)) + let image = try XCTUnwrap(context.makeImage()) + let data = NSMutableData() + let destination = try XCTUnwrap(CGImageDestinationCreateWithData( + data, + type.identifier as CFString, + 1, + nil + )) + CGImageDestinationAddImage(destination, image, nil) + XCTAssertTrue(CGImageDestinationFinalize(destination)) + return data as Data + } + } +#endif diff --git a/Tests/AriaAppleTests/Providers/FoundationModelsSessionFactoryTestSupport.swift b/Tests/AriaAppleTests/Providers/FoundationModelsSessionFactoryTestSupport.swift index d65574e..02bbfc5 100644 --- a/Tests/AriaAppleTests/Providers/FoundationModelsSessionFactoryTestSupport.swift +++ b/Tests/AriaAppleTests/Providers/FoundationModelsSessionFactoryTestSupport.swift @@ -22,7 +22,7 @@ } throw SessionFactoryTestError.expectedRequirementsReached }, - build: { _, _, _ in + build: { _, _, _, _, _ in throw SessionFactoryTestError.builderReached } ) @@ -38,18 +38,19 @@ for try await _ in stream { } XCTFail("Expected session factory validation to stop the stream", file: file, line: line) } catch let error as AgentError { - guard case let .providerFailed(_, underlying) = error else { - return XCTFail("Expected providerFailed, got \(error)", file: file, line: line) + guard case let .providerRejected(failure) = error else { + return XCTFail("Expected providerRejected, got \(error)", file: file, line: line) } + XCTAssertEqual(failure.kind, .unknown, file: file, line: line) XCTAssertEqual( - underlying?.typeName, + failure.underlying?.typeName, "SessionFactoryTestError", file: file, line: line ) XCTAssertTrue( - underlying?.message.contains("expectedRequirementsReached") == true, - "Expected the requested requirements to reach validation, got \(String(describing: underlying))", + failure.underlying?.message.contains("expectedRequirementsReached") == true, + "Expected the requested requirements to reach validation, got \(String(describing: failure.underlying))", file: file, line: line ) diff --git a/docs/layers/03-providers.md b/docs/layers/03-providers.md index fa2ee4d..29cb137 100644 --- a/docs/layers/03-providers.md +++ b/docs/layers/03-providers.md @@ -256,6 +256,30 @@ Concrete implementations live in `AriaApple/Providers/` and conform to the proto The agent layer never sees the provider's native types. +### Foundation Models multimodal prompts + +With the iOS 27 SDK and an iOS 27-family runtime, +`FoundationModelsProvider` carries image parts from the final Aria `Message` +into the native Foundation Models `Prompt`. The same conversion is used by +plain streaming and `Agent.respond(_:as:)`, so guided generation does not lose +the image while producing typed output. + +- `ImageContent.Source.data` accepts declared JPEG or PNG data, verifies the + encoded type, decodes it, and passes an Apple image attachment. Decoding also + removes encoded metadata from the prompt payload. +- `ImageContent.Source.url` and `ImageContent.Source.identifier` must be + resolved to in-memory data by the application before calling the provider. +- Invalid or unresolved sources fail as `AgentError.configurationInvalid` + before a model session is created. +- Image prompts explicitly request the Foundation Models vision capability and + fail on older runtime versions instead of silently dropping the image. + +The existing text-only path remains available on iOS 26. Consumers injecting a +custom `LanguageModel` must describe its real vision support in the supplied +`ProviderCapabilities`. Aria rejects image input when that declaration is +false, then validates the requested vision capability against the model before +generation. + ### Foundation Models model injection On iOS 27 and related Apple platform releases, `FoundationModelsProvider` can