diff --git a/README.md b/README.md index 0f5c5a8e..36115284 100644 --- a/README.md +++ b/README.md @@ -65,6 +65,38 @@ let session = LanguageModelSession(model: model, tools: [WeatherTool()]) session.toolExecutionDelegate = ToolExecutionObserver() ``` +Mirroring the OS 27 Foundation Models contract, `DynamicInstructions` can change +the instructions and tools visible to the next model request without rebuilding +the session: + +```swift +final class CurrentAppState { + var canCheckWeather = false +} + +struct CurrentAppInstructions: DynamicInstructions { + let state: CurrentAppState + + var body: some DynamicInstructions { + Instructions("Help with the currently visible app.") + if state.canCheckWeather { + WeatherTool() + } + } +} + +let state = CurrentAppState() +let session = LanguageModelSession( + model: model, + dynamicInstructions: CurrentAppInstructions(state: state), + history: savedHistory +) +``` + +The body is evaluated before every model request, including continuation +requests after tool calls. Dynamic instructions are projected into the request +context and are not persisted as the session's durable history. + ## Features ### Supported Providers diff --git a/Sources/AnyLanguageModel/DynamicInstructions.swift b/Sources/AnyLanguageModel/DynamicInstructions.swift new file mode 100644 index 00000000..02a74b7f --- /dev/null +++ b/Sources/AnyLanguageModel/DynamicInstructions.swift @@ -0,0 +1,325 @@ +/// A declarative collection of instructions and tools that a session resolves +/// immediately before each request to a language model. +/// +/// Compose values in ``body`` with ``DynamicInstructionsBuilder``. The session +/// evaluates the body again for every model request, including requests that +/// continue a response after tool execution. +@_typeEraser(AnyDynamicInstructions) +public protocol DynamicInstructions { + associatedtype Body: DynamicInstructions + + @DynamicInstructionsBuilder + var body: Body { get } +} + +/// Builds declarative dynamic instructions from instructions, tools, nested +/// dynamic instructions, and conditional content. +@resultBuilder +public struct DynamicInstructionsBuilder { + public static func buildExpression(_ expression: T) -> some DynamicInstructions where T: Tool { + DynamicTool(expression) + } + + public static func buildExpression(_ expression: T) -> T where T: DynamicInstructions { + expression + } + + public static func buildExpression(_ tools: [any Tool]) -> some DynamicInstructions { + DynamicInstructionsForEach(tools, id: \.name) { tool in + AnyDynamicInstructions(DynamicTool(tool)) + } + } + + @_disfavoredOverload + public static func buildBlock( + _ contents: repeat each Content + ) -> TupleDynamicInstructions + where repeat each Content: DynamicInstructions { + TupleDynamicInstructions(repeat each contents) + } + + public static func buildBlock(_ content: T) -> T where T: DynamicInstructions { + content + } + + public static func buildBlock() -> EmptyDynamicInstructions { + EmptyDynamicInstructions() + } + + public static func buildEither( + first content: TrueContent + ) -> ConditionalDynamicInstructions + where TrueContent: DynamicInstructions, FalseContent: DynamicInstructions { + ConditionalDynamicInstructions(.trueContent(content)) + } + + public static func buildEither( + second content: FalseContent + ) -> ConditionalDynamicInstructions + where TrueContent: DynamicInstructions, FalseContent: DynamicInstructions { + ConditionalDynamicInstructions(.falseContent(content)) + } + + public static func buildOptional(_ content: Content?) -> Content? + where Content: DynamicInstructions { + content + } + + public static func buildLimitedAvailability( + _ content: some DynamicInstructions + ) -> AnyDynamicInstructions { + AnyDynamicInstructions(content) + } +} + +/// A type-erased dynamic-instructions value. +public struct AnyDynamicInstructions: DynamicInstructions { + public typealias Body = Never + + fileprivate let resolveValue: () -> ResolvedDynamicInstructions + + public init(_ dynamicInstructions: any DynamicInstructions) { + resolveValue = { resolveDynamicInstructions(dynamicInstructions) } + } + + public init(erasing dynamicInstructions: some DynamicInstructions) { + self.init(dynamicInstructions) + } + + public var body: Never { + fatalError("AnyDynamicInstructions has no body") + } + + func resolveForRequest() -> ResolvedDynamicInstructions { + resolveValue() + } +} + +/// A dynamic-instructions value that contains an ordered tuple of components. +public struct TupleDynamicInstructions: DynamicInstructions +where repeat each Content: DynamicInstructions { + public typealias Body = Never + + fileprivate let contents: (repeat each Content) + + public init(_ contents: repeat each Content) { + self.contents = (repeat each contents) + } + + public var body: Never { + fatalError("TupleDynamicInstructions has no body") + } +} + +/// A dynamic-instructions value that contains one of two branches. +public struct ConditionalDynamicInstructions: DynamicInstructions +where TrueContent: DynamicInstructions, FalseContent: DynamicInstructions { + public enum Branch { + case trueContent(TrueContent) + case falseContent(FalseContent) + } + + public typealias Body = Never + + fileprivate let branch: Branch + + public init(_ branch: Branch) { + self.branch = branch + } + + public var body: Never { + fatalError("ConditionalDynamicInstructions has no body") + } +} + +extension Optional: DynamicInstructions where Wrapped: DynamicInstructions { + public typealias Body = Never + + public var body: Never { + fatalError("Optional dynamic instructions have no body") + } +} + +extension Never: DynamicInstructions { + public typealias Body = Never + + public var body: Never { self } +} + +/// An empty dynamic-instructions value. +public struct EmptyDynamicInstructions: DynamicInstructions, Sendable { + public typealias Body = Never + + public init() {} + + public var body: Never { + fatalError("EmptyDynamicInstructions has no body") + } +} + +/// Builds dynamic instructions from a collection. +public struct DynamicInstructionsForEach: DynamicInstructions +where Data: RandomAccessCollection, ID: Hashable, Content: DynamicInstructions { + public typealias Body = Never + + fileprivate let data: Data + fileprivate let id: KeyPath + fileprivate let content: (Data.Element) -> Content + + public init( + _ data: Data, + id: KeyPath, + @DynamicInstructionsBuilder content: @escaping (Data.Element) -> Content + ) { + self.data = data + self.id = id + self.content = content + } + + public var body: Never { + fatalError("DynamicInstructionsForEach has no body") + } +} + +extension DynamicInstructionsForEach where ID == Data.Element.ID, Data.Element: Identifiable { + public init( + _ data: Data, + @DynamicInstructionsBuilder content: @escaping (Data.Element) -> Content + ) { + self.init(data, id: \.id, content: content) + } +} + +extension DynamicInstructions { + public typealias ForEach = DynamicInstructionsForEach +} + +extension Instructions: DynamicInstructions { + public var body: some DynamicInstructions { + EmptyDynamicInstructions() + } +} + +struct ResolvedDynamicInstructions: Sendable { + let instructions: Instructions? + let tools: [any Tool] + + fileprivate init(instructions: Instructions?, tools: [any Tool]) { + self.instructions = instructions + self.tools = tools + } + + fileprivate static let empty = Self(instructions: nil, tools: []) + + fileprivate func appending(_ other: Self) -> Self { + let combinedInstructions: Instructions? + switch (instructions, other.instructions) { + case (nil, nil): + combinedInstructions = nil + case (let instructions?, nil), (nil, let instructions?): + combinedInstructions = instructions + case (let first?, let second?): + combinedInstructions = Instructions { + first + second + } + } + return Self( + instructions: combinedInstructions, + tools: tools + other.tools + ) + } +} + +private protocol PrimitiveDynamicInstructions { + func resolve() -> ResolvedDynamicInstructions +} + +private struct DynamicTool: DynamicInstructions, PrimitiveDynamicInstructions { + typealias Body = Never + + let tool: any Tool + + init(_ tool: any Tool) { + self.tool = tool + } + + var body: Never { + fatalError("DynamicTool has no body") + } + + func resolve() -> ResolvedDynamicInstructions { + ResolvedDynamicInstructions(instructions: nil, tools: [tool]) + } +} + +extension AnyDynamicInstructions: PrimitiveDynamicInstructions { + fileprivate func resolve() -> ResolvedDynamicInstructions { + resolveValue() + } +} + +extension TupleDynamicInstructions: PrimitiveDynamicInstructions { + fileprivate func resolve() -> ResolvedDynamicInstructions { + var result = ResolvedDynamicInstructions.empty + repeat result = result.appending(resolveDynamicInstructions(each contents)) + return result + } +} + +extension ConditionalDynamicInstructions: PrimitiveDynamicInstructions { + fileprivate func resolve() -> ResolvedDynamicInstructions { + switch branch { + case .trueContent(let content): + resolveDynamicInstructions(content) + case .falseContent(let content): + resolveDynamicInstructions(content) + } + } +} + +extension Optional: PrimitiveDynamicInstructions where Wrapped: DynamicInstructions { + fileprivate func resolve() -> ResolvedDynamicInstructions { + map(resolveDynamicInstructions) ?? .empty + } +} + +extension Never: PrimitiveDynamicInstructions { + fileprivate func resolve() -> ResolvedDynamicInstructions { + switch self {} + } +} + +extension EmptyDynamicInstructions: PrimitiveDynamicInstructions { + fileprivate func resolve() -> ResolvedDynamicInstructions { + .empty + } +} + +extension DynamicInstructionsForEach: PrimitiveDynamicInstructions { + fileprivate func resolve() -> ResolvedDynamicInstructions { + data.reduce(into: .empty) { result, element in + result = result.appending(resolveDynamicInstructions(content(element))) + } + } +} + +extension Instructions: PrimitiveDynamicInstructions { + fileprivate func resolve() -> ResolvedDynamicInstructions { + ResolvedDynamicInstructions(instructions: self, tools: []) + } +} + +private func resolveDynamicInstructions( + _ dynamicInstructions: any DynamicInstructions +) -> ResolvedDynamicInstructions { + func resolve(_ content: Content) -> ResolvedDynamicInstructions + where Content: DynamicInstructions { + if let primitive = content as? any PrimitiveDynamicInstructions { + return primitive.resolve() + } + return resolve(content.body) + } + + return resolve(dynamicInstructions) +} diff --git a/Sources/AnyLanguageModel/LanguageModelSession.swift b/Sources/AnyLanguageModel/LanguageModelSession.swift index f027b894..00a971b6 100644 --- a/Sources/AnyLanguageModel/LanguageModelSession.swift +++ b/Sources/AnyLanguageModel/LanguageModelSession.swift @@ -66,6 +66,34 @@ public final class LanguageModelSession: @unchecked Sendable { private let model: any LanguageModel public let tools: [any Tool] public let instructions: Instructions? + private let dynamicInstructions: AnyDynamicInstructions? + private let dynamicInstructionsLock = NSLock() + + nonisolated var usesDynamicInstructions: Bool { + dynamicInstructions != nil + } + + /// The immutable session inputs resolved for one model request. + /// + /// Language-model implementations should create one context immediately + /// before each provider or local-model request. If that request produces + /// tool calls, execute them with ``tools`` from the same context. Resolve a + /// new context only before the continuation request. + public struct RequestContext: Sendable { + public let transcript: Transcript + public let instructions: Instructions? + public let tools: [any Tool] + + fileprivate init( + transcript: Transcript, + instructions: Instructions?, + tools: [any Tool] + ) { + self.transcript = transcript + self.instructions = instructions + self.tools = tools + } + } /// A delegate that observes and controls tool execution. /// @@ -90,7 +118,13 @@ public final class LanguageModelSession: @unchecked Sendable { tools: [any Tool] = [], instructions: String ) { - self.init(model: model, tools: tools, instructions: Instructions(instructions), transcript: Transcript()) + self.init( + model: model, + tools: tools, + instructions: Instructions(instructions), + dynamicInstructions: nil, + transcript: Transcript() + ) } public convenience init( @@ -98,7 +132,13 @@ public final class LanguageModelSession: @unchecked Sendable { tools: [any Tool] = [], instructions: Instructions? = nil ) { - self.init(model: model, tools: tools, instructions: instructions, transcript: Transcript()) + self.init( + model: model, + tools: tools, + instructions: instructions, + dynamicInstructions: nil, + transcript: Transcript() + ) } public convenience init( @@ -106,17 +146,45 @@ public final class LanguageModelSession: @unchecked Sendable { tools: [any Tool] = [], transcript: Transcript ) { - self.init(model: model, tools: tools, instructions: nil, transcript: transcript) + self.init( + model: model, + tools: tools, + instructions: nil, + dynamicInstructions: nil, + transcript: transcript + ) + } + + /// Creates a session whose instructions and tools are resolved before each + /// model request. + /// + /// The history excludes the dynamic instructions entry. Resolved dynamic + /// instructions are projected only into ``RequestContext/transcript`` and + /// never become durable transcript state. + public convenience init( + model: any LanguageModel, + dynamicInstructions: sending some DynamicInstructions, + history: some Collection = [] + ) { + self.init( + model: model, + tools: [], + instructions: nil, + dynamicInstructions: AnyDynamicInstructions(dynamicInstructions), + transcript: Transcript(entries: Array(history)) + ) } private init( model: any LanguageModel, tools: [any Tool], instructions: Instructions?, + dynamicInstructions: AnyDynamicInstructions?, transcript: Transcript ) { self.model = model self.tools = tools + self.dynamicInstructions = dynamicInstructions let resolvedInstructions = instructions ?? Self.instructions(from: transcript) self.instructions = resolvedInstructions @@ -146,6 +214,44 @@ public final class LanguageModelSession: @unchecked Sendable { self.state = .init(.init(finalTranscript)) } + /// Resolves the instructions, tools, and transcript view for the next model + /// request. + /// + /// This is an AnyLanguageModel provider-integration seam. Calling it does + /// not mutate the session transcript or any global tool registry. + nonisolated public func resolvedRequestContext() -> RequestContext { + guard let dynamicInstructions else { + return RequestContext( + transcript: transcript, + instructions: instructions, + tools: tools + ) + } + + let resolved = dynamicInstructionsLock.withLock { + dynamicInstructions.resolveForRequest() + } + var requestTranscript = transcript + if let instructions = resolved.instructions { + let instructionsEntry = Transcript.Entry.instructions( + Transcript.Instructions( + segments: [ + .text(Transcript.TextSegment(content: instructions.description)) + ], + toolDefinitions: resolved.tools + .filter(\.includesSchemaInInstructions) + .map { Transcript.ToolDefinition(tool: $0) } + ) + ) + requestTranscript = Transcript(entries: [instructionsEntry] + requestTranscript) + } + return RequestContext( + transcript: requestTranscript, + instructions: resolved.instructions, + tools: resolved.tools + ) + } + private static func instructions(from transcript: Transcript) -> Instructions? { guard case .instructions(let instructions)? = transcript.first else { return nil } guard instructions.segments.count == 1, diff --git a/Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift b/Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift index 2f7dd7cc..1224ed7e 100644 --- a/Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift @@ -442,9 +442,10 @@ public struct AnthropicLanguageModel: LanguageModel { ) async throws -> LanguageModelSession.Response where Content: Generable { let url = baseURL.appendingPathComponent("v1/messages") let headers = buildHeaders() + let requestContext = session.resolvedRequestContext() // Convert available tools to Anthropic format - let anthropicTools: [AnthropicTool] = try session.tools.map { tool in + let anthropicTools: [AnthropicTool] = try requestContext.tools.map { tool in try convertToolToAnthropicFormat(tool) } @@ -452,7 +453,7 @@ public struct AnthropicLanguageModel: LanguageModel { let params = try createMessageParams( model: model, system: nil, - messages: session.transcript.toAnthropicMessages(), + messages: requestContext.transcript.toAnthropicMessages(), tools: anthropicTools.isEmpty ? nil : anthropicTools, responseSchema: responseSchema, options: options @@ -477,7 +478,11 @@ public struct AnthropicLanguageModel: LanguageModel { } if !toolUses.isEmpty { - let resolution = try await resolveToolUses(toolUses, session: session) + let resolution = try await resolveToolUses( + toolUses, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -575,19 +580,18 @@ public struct AnthropicLanguageModel: LanguageModel { let task = Task { @Sendable in do { let headers = buildHeaders() - - // Convert available tools to Anthropic format - let anthropicTools: [AnthropicTool] = try session.tools.map { tool in - try convertToolToAnthropicFormat(tool) - } - let responseSchema = type == String.self ? nil : try convertSchemaToAnthropicFormat(schema) - var messages = session.transcript.toAnthropicMessages() + var inFlightMessages: [AnthropicMessage] = [] var state = StreamingResponseState() var toolRounds = ToolRoundLimit(provider: "Anthropic") while true { try Task.checkCancellation() + let requestContext = session.resolvedRequestContext() + let anthropicTools: [AnthropicTool] = try requestContext.tools.map { + try convertToolToAnthropicFormat($0) + } + let messages = requestContext.transcript.toAnthropicMessages() + inFlightMessages var params = try createMessageParams( model: model, system: nil, @@ -667,14 +671,18 @@ public struct AnthropicLanguageModel: LanguageModel { guard !toolUses.isEmpty else { break } try Task.checkCancellation() try toolRounds.record(toolUses.map(\.roundCall)) - switch try await resolveToolUses(toolUses, session: session) { + switch try await resolveToolUses( + toolUses, + tools: requestContext.tools, + session: session + ) { case .stop(let calls): state.entries.append(.toolCalls(Transcript.ToolCalls(calls))) continuation.yield(try state.stoppedSnapshot()) continuation.finish() return case .invocations(let invocations): - messages.append(.init(role: .assistant, content: content)) + inFlightMessages.append(.init(role: .assistant, content: content)) state.entries.append(.toolCalls(Transcript.ToolCalls(invocations.map(\.call)))) var results: [AnthropicContent] = [] for invocation in invocations { @@ -688,7 +696,7 @@ public struct AnthropicLanguageModel: LanguageModel { ) ) } - messages.append(.init(role: .user, content: results)) + inFlightMessages.append(.init(role: .user, content: results)) } if let snapshot = snapshot() { continuation.yield(snapshot) } state.beginNextRound() @@ -868,12 +876,13 @@ private func convertSchemaToAnthropicFormat(_ schema: GenerationSchema) throws - private func resolveToolUses( _ toolUses: [AnthropicToolUse], + tools: [any Tool], session: LanguageModelSession ) async throws -> ToolResolutionOutcome { if toolUses.isEmpty { return .invocations([]) } var toolsByName: [String: any Tool] = [:] - for tool in session.tools { + for tool in tools { if toolsByName[tool.name] == nil { toolsByName[tool.name] = tool } diff --git a/Sources/AnyLanguageModel/Models/CoreMLLanguageModel.swift b/Sources/AnyLanguageModel/Models/CoreMLLanguageModel.swift index 05c3aebd..9c3514ef 100644 --- a/Sources/AnyLanguageModel/Models/CoreMLLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/CoreMLLanguageModel.swift @@ -113,11 +113,12 @@ includeSchemaInPrompt: Bool, options: GenerationOptions ) async throws -> LanguageModelSession.Response where Content: Generable { - try validateNoImageSegments(in: session) + let requestContext = session.resolvedRequestContext() + try validateNoImageSegments(in: requestContext.transcript) if type != String.self { let (jsonString, usage) = try await generateStructuredJSON( - session: session, + requestContext: requestContext, prompt: prompt, schema: schema, options: options, @@ -139,8 +140,8 @@ let tokens: [Int] if let chatTemplateHandler = chatTemplateHandler { // Use chat template handler with optional tools - let messages = chatTemplateHandler(session.instructions, prompt) - let toolSpecs: [ToolSpec]? = toolsHandler?(session.tools) + let messages = chatTemplateHandler(requestContext.instructions, prompt) + let toolSpecs: [ToolSpec]? = toolsHandler?(requestContext.tools) tokens = try tokenizer.applyChatTemplate(messages: messages, tools: toolSpecs) } else { // Fall back to direct tokenizer encoding @@ -227,17 +228,6 @@ } } - // Validate that no image segments are present - do { - try validateNoImageSegments(in: session) - } catch { - return LanguageModelSession.ResponseStream( - stream: AsyncThrowingStream { continuation in - continuation.finish(throwing: error) - } - ) - } - // Convert AnyLanguageModel GenerationOptions to swift-transformers GenerationConfig let generationConfig = toGenerationConfig(options) @@ -246,11 +236,13 @@ @Sendable continuation in let task = Task { do { + let requestContext = session.resolvedRequestContext() + try validateNoImageSegments(in: requestContext.transcript) let tokens: [Int] if let chatTemplateHandler = chatTemplateHandler { // Use chat template handler with optional tools - let messages = chatTemplateHandler(session.instructions, prompt) - let toolSpecs: [ToolSpec]? = toolsHandler?(session.tools) + let messages = chatTemplateHandler(requestContext.instructions, prompt) + let toolSpecs: [ToolSpec]? = toolsHandler?(requestContext.tools) tokens = try tokenizer.applyChatTemplate(messages: messages, tools: toolSpecs) } else { // Fall back to direct tokenizer encoding @@ -298,10 +290,10 @@ // MARK: - Image Validation - private func validateNoImageSegments(in session: LanguageModelSession) throws { + private func validateNoImageSegments(in transcript: Transcript) throws { // Note: Instructions is a plain text type without segments, so no image check needed there. // Check for image segments in the most recent prompt - for entry in session.transcript.reversed() { + for entry in transcript.reversed() { if case .prompt(let p) = entry { for segment in p.segments { if case .image = segment { @@ -406,7 +398,7 @@ } private func generateStructuredJSON( - session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, prompt: Prompt, schema: GenerationSchema, options: GenerationOptions, @@ -416,7 +408,7 @@ var generationConfig = toStructuredGenerationConfig(options) let promptTokens = try structuredPromptTokens( - in: session, + requestContext: requestContext, prompt: prompt, schema: schema, includeSchemaInPrompt: includeSchemaInPrompt @@ -453,20 +445,20 @@ } private func structuredPromptTokens( - in session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, prompt: Prompt, schema: GenerationSchema, includeSchemaInPrompt: Bool ) throws -> [Int] { if let chatTemplateHandler = chatTemplateHandler { - var messages = chatTemplateHandler(session.instructions, prompt) + var messages = chatTemplateHandler(requestContext.instructions, prompt) if includeSchemaInPrompt { let schemaPrompt = schemaPrompt(for: schema) if !schemaPrompt.isEmpty { messages.insert(["role": "system", "content": schemaPrompt], at: 0) } } - let toolSpecs: [ToolSpec]? = toolsHandler?(session.tools) + let toolSpecs: [ToolSpec]? = toolsHandler?(requestContext.tools) return try tokenizer.applyChatTemplate(messages: messages, tools: toolSpecs) } diff --git a/Sources/AnyLanguageModel/Models/FoundationLanguageModel.swift b/Sources/AnyLanguageModel/Models/FoundationLanguageModel.swift index 307c992d..b5bf8599 100644 --- a/Sources/AnyLanguageModel/Models/FoundationLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/FoundationLanguageModel.swift @@ -105,13 +105,13 @@ } private func makeSession( - tools: [any FoundationModels.Tool], - transcript: FoundationModels.Transcript + for session: LanguageModelSession, + prompt: Prompt ) async throws -> FoundationModels.LanguageModelSession { - FoundationModels.LanguageModelSession( + makeFoundationModelsSession( model: try await loadedModel(), - tools: tools, - transcript: transcript + session: session, + prompt: prompt ) } @@ -157,16 +157,8 @@ includeSchemaInPrompt: Bool, options: GenerationOptions ) async throws -> LanguageModelSession.Response where Content: Generable { - let fmTools = session.tools.toFoundationModels() - let fmTranscript = fmTranscriptDroppingDuplicatePrompt(session.transcript, prompt: prompt) - .toFoundationModels( - instructions: session.instructions, - toolDefinitions: session.tools - .filter(\.includesSchemaInInstructions) - .map { Transcript.ToolDefinition(tool: $0) } - ) return try await fmRespond( - makeSession: { try await self.makeSession(tools: fmTools, transcript: fmTranscript) }, + makeSession: { try await self.makeSession(for: session, prompt: prompt) }, fmPrompt: prompt.toFoundationModels(), fmOptions: options.toFoundationModels(), type: type, @@ -217,16 +209,8 @@ includeSchemaInPrompt: Bool, options: GenerationOptions ) -> sending LanguageModelSession.ResponseStream where Content: Generable { - let fmTools = session.tools.toFoundationModels() - let fmTranscript = fmTranscriptDroppingDuplicatePrompt(session.transcript, prompt: prompt) - .toFoundationModels( - instructions: session.instructions, - toolDefinitions: session.tools - .filter(\.includesSchemaInInstructions) - .map { Transcript.ToolDefinition(tool: $0) } - ) return fmStreamResponse( - makeSession: { try await self.makeSession(tools: fmTools, transcript: fmTranscript) }, + makeSession: { try await self.makeSession(for: session, prompt: prompt) }, fmPrompt: prompt.toFoundationModels(), fmOptions: options.toFoundationModels(), type: type, diff --git a/Sources/AnyLanguageModel/Models/GeminiLanguageModel.swift b/Sources/AnyLanguageModel/Models/GeminiLanguageModel.swift index ba06a0c8..1035505e 100644 --- a/Sources/AnyLanguageModel/Models/GeminiLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/GeminiLanguageModel.swift @@ -313,12 +313,10 @@ public struct GeminiLanguageModel: LanguageModel { .appendingPathComponent("models/\(model):generateContent") let headers = buildHeaders() - let geminiTools = try buildTools(from: session.tools, serverTools: effectiveServerTools) + var inFlightEntries: [Transcript.Entry] = [] - var transcript = session.transcript - - // The entries this call adds, which is what the response reports. `transcript` keeps the - // full conversation because each iteration rebuilds the request from it. + // The entries this call adds, which is what the response reports. `inFlightEntries` + // preserves tool rounds while each iteration rebuilds the request from a fresh context. var entries: [Transcript.Entry] = [] var usage = ReportedUsage() // The text of earlier tool rounds, which string responses include. @@ -327,8 +325,17 @@ public struct GeminiLanguageModel: LanguageModel { var toolRounds = ToolRoundLimit(provider: "Gemini") // Multi-turn conversation loop for tool calling while true { + let requestContext = session.resolvedRequestContext() + let geminiTools = try buildTools( + from: requestContext.tools, + serverTools: effectiveServerTools + ) + var requestTranscript = requestContext.transcript + for entry in inFlightEntries { + requestTranscript.append(entry) + } let params = try createGenerateContentParams( - contents: transcript.toGeminiContent(), + contents: requestTranscript.toGeminiContent(), tools: geminiTools, generating: type, schema: schema, @@ -365,7 +372,11 @@ public struct GeminiLanguageModel: LanguageModel { if !functionCalls.isEmpty { // Resolve function calls try toolRounds.record(functionCalls.map(\.roundCall)) - let resolution = try await resolveFunctionCalls(functionCalls, session: session) + let resolution = try await resolveFunctionCalls( + functionCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -383,12 +394,12 @@ public struct GeminiLanguageModel: LanguageModel { let calls = Transcript.Entry.toolCalls( Transcript.ToolCalls(invocations.map(\.call), providerMetadata: providerMetadata) ) - transcript.append(calls) + inFlightEntries.append(calls) entries.append(calls) for invocation in invocations { let output = Transcript.Entry.toolOutput(invocation.output) - transcript.append(output) + inFlightEntries.append(output) entries.append(output) } } @@ -485,16 +496,22 @@ public struct GeminiLanguageModel: LanguageModel { let task = Task { @Sendable in do { let headers = buildHeaders() - - let geminiTools = try buildTools(from: session.tools, serverTools: effectiveServerTools) - - var transcript = session.transcript + var inFlightEntries: [Transcript.Entry] = [] var state = StreamingResponseState() var toolRounds = ToolRoundLimit(provider: "Gemini") while true { try Task.checkCancellation() + let requestContext = session.resolvedRequestContext() + let geminiTools = try buildTools( + from: requestContext.tools, + serverTools: effectiveServerTools + ) + var requestTranscript = requestContext.transcript + for entry in inFlightEntries { + requestTranscript.append(entry) + } let params = try createGenerateContentParams( - contents: transcript.toGeminiContent(), + contents: requestTranscript.toGeminiContent(), tools: geminiTools, generating: type, schema: schema, @@ -535,7 +552,11 @@ public struct GeminiLanguageModel: LanguageModel { try Task.checkCancellation() let metadata = try textPartMetadata(parts, includeUnsignedText: true) try toolRounds.record(functionCalls.map(\.roundCall)) - switch try await resolveFunctionCalls(functionCalls, session: session) { + switch try await resolveFunctionCalls( + functionCalls, + tools: requestContext.tools, + session: session + ) { case .stop(let calls): state.entries.append(.toolCalls(Transcript.ToolCalls(calls, providerMetadata: metadata))) continuation.yield(try state.stoppedSnapshot()) @@ -545,11 +566,11 @@ public struct GeminiLanguageModel: LanguageModel { let calls = Transcript.Entry.toolCalls( Transcript.ToolCalls(invocations.map(\.call), providerMetadata: metadata) ) - transcript.append(calls) + inFlightEntries.append(calls) state.entries.append(calls) for invocation in invocations { let output = Transcript.Entry.toolOutput(invocation.output) - transcript.append(output) + inFlightEntries.append(output) state.entries.append(output) } } @@ -696,12 +717,13 @@ private enum ToolResolutionOutcome { private func resolveFunctionCalls( _ functionCalls: [GeminiFunctionCall], + tools: [any Tool], session: LanguageModelSession ) async throws -> ToolResolutionOutcome { if functionCalls.isEmpty { return .invocations([]) } var toolsByName: [String: any Tool] = [:] - for tool in session.tools { + for tool in tools { if toolsByName[tool.name] == nil { toolsByName[tool.name] = tool } diff --git a/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift b/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift index 68c5ea1d..c6b560bf 100644 --- a/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift @@ -843,9 +843,14 @@ import Foundation return LlamaToolCallFormat.detect(template: template) } - private func makeToolPromptContext(for session: LanguageModelSession) throws -> LlamaToolPromptContext? { - guard !session.tools.isEmpty, self.model != nil else { return nil } - return try LlamaToolPromptContext(format: currentToolCallFormat(), tools: session.tools) + private func makeToolPromptContext( + tools: [any Tool], + pendingEntries: [Transcript.Entry] = [] + ) throws -> LlamaToolPromptContext? { + guard (!tools.isEmpty || !pendingEntries.isEmpty), self.model != nil else { return nil } + var context = try LlamaToolPromptContext(format: currentToolCallFormat(), tools: tools) + context.pendingEntries = pendingEntries + return context } private func makeTranscriptToolCalls( @@ -862,12 +867,13 @@ import Foundation private func resolveToolCalls( _ parsedCalls: [LlamaParsedToolCall], + tools: [any Tool], session: LanguageModelSession ) async throws -> ToolResolutionOutcome { if parsedCalls.isEmpty { return .invocations([]) } var toolsByName: [String: any Tool] = [:] - for tool in session.tools where toolsByName[tool.name] == nil { + for tool in tools where toolsByName[tool.name] == nil { toolsByName[tool.name] = tool } @@ -989,9 +995,6 @@ import Foundation includeSchemaInPrompt: Bool, options: GenerationOptions ) async throws -> LanguageModelSession.Response where Content: Generable { - if mmprojPath == nil { - try validateNoImageSegments(in: session) - } try ensureModelLoaded() let runtimeOptions = resolvedOptions(from: options) @@ -1002,18 +1005,26 @@ import Foundation if type == String.self { let maxTokens = runtimeOptions.maximumResponseTokens ?? 100 let outputFormat = currentToolCallFormat() - var toolContext = try makeToolPromptContext(for: session) let maxToolIterations = 8 var toolIteration = 0 var previousToolCallSignature: String? var allEntries: [Transcript.Entry] = [] var text = "" var usage = LanguageModelSession.Usage.zero + var pendingEntries: [Transcript.Entry] = [] generationLoop: while true { + let requestContext = session.resolvedRequestContext() + if mmprojPath == nil { + try validateNoImageSegments(in: requestContext.transcript) + } + let toolContext = try makeToolPromptContext( + tools: requestContext.tools, + pendingEntries: pendingEntries + ) var promptImages: [Data] = [] let fullPrompt = try formatPrompt( - for: session, + requestContext: requestContext, extraSystemMessage: nil, assistantPrefill: runtimeOptions.assistantPrefill, imageMarker: imageMarker, @@ -1091,7 +1102,11 @@ import Foundation } previousToolCallSignature = signature - let resolution = try await resolveToolCalls(parsedCalls, session: session) + let resolution = try await resolveToolCalls( + parsedCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -1111,11 +1126,11 @@ import Foundation Transcript.ToolCalls(invocations.map(\.call)) ) allEntries.append(callsEntry) - toolContext?.pendingEntries.append(callsEntry) + pendingEntries.append(callsEntry) for invocation in invocations { let outputEntry = Transcript.Entry.toolOutput(invocation.output) allEntries.append(outputEntry) - toolContext?.pendingEntries.append(outputEntry) + pendingEntries.append(outputEntry) } } } @@ -1127,11 +1142,15 @@ import Foundation usage: usage ) } else { + let requestContext = session.resolvedRequestContext() + if mmprojPath == nil { + try validateNoImageSegments(in: requestContext.transcript) + } var promptImages: [Data] = [] let fullPrompt: String if includeSchemaInPrompt { fullPrompt = try formatPrompt( - for: session, + requestContext: requestContext, extraSystemMessage: schemaPrompt(for: schema), assistantPrefill: runtimeOptions.assistantPrefill, imageMarker: imageMarker, @@ -1139,7 +1158,7 @@ import Foundation ) } else { fullPrompt = try formatPrompt( - for: session, + requestContext: requestContext, extraSystemMessage: nil, assistantPrefill: runtimeOptions.assistantPrefill, imageMarker: imageMarker, @@ -1227,18 +1246,6 @@ import Foundation } } - if mmprojPath == nil { - do { - try validateNoImageSegments(in: session) - } catch { - return LanguageModelSession.ResponseStream( - stream: AsyncThrowingStream { continuation in - continuation.finish(throwing: error) - } - ) - } - } - let stream: AsyncThrowingStream.Snapshot, any Error> = AsyncThrowingStream { continuation in let task = Task { @@ -1248,7 +1255,6 @@ import Foundation let runtimeOptions = resolvedOptions(from: options) let maxTokens = runtimeOptions.maximumResponseTokens ?? 100 let outputFormat = self.currentToolCallFormat() - var toolContext = try self.makeToolPromptContext(for: session) let maxToolIterations = 8 var toolIteration = 0 var previousToolCallSignature: String? @@ -1256,6 +1262,7 @@ import Foundation var emittedBase = "" var usage = LanguageModelSession.Usage.zero var lastYieldedText: String? + var pendingEntries: [Transcript.Entry] = [] let imageMarker = self.mtmdContext != nil ? String(cString: mtmd_default_marker()) : nil @@ -1271,9 +1278,17 @@ import Foundation } generationLoop: while true { + let requestContext = session.resolvedRequestContext() + if self.mmprojPath == nil { + try self.validateNoImageSegments(in: requestContext.transcript) + } + let toolContext = try self.makeToolPromptContext( + tools: requestContext.tools, + pendingEntries: pendingEntries + ) var promptImages: [Data] = [] let fullPrompt = try self.formatPrompt( - for: session, + requestContext: requestContext, extraSystemMessage: nil, assistantPrefill: runtimeOptions.assistantPrefill, imageMarker: imageMarker, @@ -1366,7 +1381,11 @@ import Foundation } previousToolCallSignature = signature - let resolution = try await self.resolveToolCalls(parsedCalls, session: session) + let resolution = try await self.resolveToolCalls( + parsedCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): emittedBase += roundVisible @@ -1384,11 +1403,11 @@ import Foundation Transcript.ToolCalls(invocations.map(\.call)) ) accumulatedEntries.append(callsEntry) - toolContext?.pendingEntries.append(callsEntry) + pendingEntries.append(callsEntry) for invocation in invocations { let outputEntry = Transcript.Entry.toolOutput(invocation.output) accumulatedEntries.append(outputEntry) - toolContext?.pendingEntries.append(outputEntry) + pendingEntries.append(outputEntry) } emittedBase += roundVisible yieldSnapshot(emittedBase) @@ -2133,9 +2152,9 @@ import Foundation // MARK: - Image Validation - private func validateNoImageSegments(in session: LanguageModelSession) throws { + private func validateNoImageSegments(in transcript: Transcript) throws { // Check for image segments in the most recent prompt from the transcript - for entry in session.transcript.reversed() { + for entry in transcript.reversed() { if case .prompt(let p) = entry { for segment in p.segments { if case .image = segment { @@ -2268,14 +2287,14 @@ import Foundation } private func formatPrompt( - for session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, extraSystemMessage: String? = nil, assistantPrefill: String? = nil, toolContext: LlamaToolPromptContext? = nil ) throws -> String { var images: [Data] = [] return try formatPrompt( - for: session, + requestContext: requestContext, extraSystemMessage: extraSystemMessage, assistantPrefill: assistantPrefill, imageMarker: nil, @@ -2285,7 +2304,7 @@ import Foundation } private func formatPrompt( - for session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, extraSystemMessage: String?, assistantPrefill: String?, imageMarker: String?, @@ -2359,7 +2378,7 @@ import Foundation } } - for entry in session.transcript { + for entry in requestContext.transcript { try appendEntry(entry) } if let toolContext { diff --git a/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift b/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift index c8a2d780..d86e7384 100644 --- a/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift @@ -918,8 +918,8 @@ import Foundation GPUMemoryManager.shared.markIdle(scope: id) } - private func mlxToolSpecs(for session: LanguageModelSession) -> [ToolSpec]? { - session.tools.isEmpty ? nil : session.tools.map { convertToolToMLXSpec($0) } + private func mlxToolSpecs(for tools: [any Tool]) -> [ToolSpec]? { + tools.isEmpty ? nil : tools.map { convertToolToMLXSpec($0) } } private func makeUserInput( @@ -991,10 +991,11 @@ import Foundation defer { endGenerationScope(generationScope) } if type != String.self { + let requestContext = session.resolvedRequestContext() let (jsonString, usage) = try await generateStructuredJSON( context: context, tokenCache: loaded.tokenCache, - session: session, + requestContext: requestContext, prompt: prompt, schema: schema, options: options, @@ -1010,8 +1011,6 @@ import Foundation ) } - let toolSpecs = mlxToolSpecs(for: session) - // Map AnyLanguageModel GenerationOptions to MLX GenerateParameters let generateParameters = toGenerateParameters(options) @@ -1021,8 +1020,7 @@ import Foundation options[custom: MLXLanguageModel.self]?.processingForUserInput ?? .init(resize: nil) - // Build chat history from full transcript - var chat = convertTranscriptToMLXChat(session: session, fallbackPrompt: prompt.description) + var pendingChat: [MLXLMCommon.Chat.Message] = [] var usage = LanguageModelSession.Usage.zero var allTextChunks: [String] = [] @@ -1033,6 +1031,13 @@ import Foundation // Loop until no more tool calls while true { + let requestContext = session.resolvedRequestContext() + let toolSpecs = mlxToolSpecs(for: requestContext.tools) + let chat = + convertTranscriptToMLXChat( + requestContext: requestContext, + fallbackPrompt: prompt.description + ) + pendingChat // Build user input with current chat history and tools let userInput = makeUserInput( chat: chat, @@ -1087,7 +1092,7 @@ import Foundation // Add assistant response to chat history if !assistantText.isEmpty { - chat.append(.assistant(assistantText)) + pendingChat.append(.assistant(assistantText)) } // If there are tool calls, execute them and continue @@ -1110,7 +1115,11 @@ import Foundation } previousToolCallSignature = signature - let resolution = try await resolveToolCalls(collectedToolCalls, session: session) + let resolution = try await resolveToolCalls( + collectedToolCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -1132,7 +1141,7 @@ import Foundation // Convert tool output to JSON string for MLX let toolResultJSON = toolOutputToJSON(invocation.output) - chat.append(.tool(toolResultJSON)) + pendingChat.append(.tool(toolResultJSON)) } // Continue loop to generate with tool results @@ -1259,11 +1268,7 @@ import Foundation let userInputProcessing = options[custom: MLXLanguageModel.self]?.processingForUserInput ?? .init(resize: nil) - let toolSpecs = mlxToolSpecs(for: session) - var chat = convertTranscriptToMLXChat( - session: session, - fallbackPrompt: prompt.description - ) + var pendingChat: [MLXLMCommon.Chat.Message] = [] // Accumulators live outside the tool loop so streamed snapshots stay // monotonic across rounds: text never shrinks, entries only grow. @@ -1291,6 +1296,13 @@ import Foundation // Loop until the model stops without pending tool calls (mirrors `respond()`). toolLoop: while true { + let requestContext = session.resolvedRequestContext() + let toolSpecs = mlxToolSpecs(for: requestContext.tools) + let chat = + convertTranscriptToMLXChat( + requestContext: requestContext, + fallbackPrompt: prompt.description + ) + pendingChat let userInput = makeUserInput( chat: chat, tools: toolSpecs, @@ -1346,7 +1358,7 @@ import Foundation // Feed this round's assistant text back into the chat history. let roundText = String(accumulatedText.dropFirst(roundStartTextCount)) if !roundText.isEmpty { - chat.append(.assistant(roundText)) + pendingChat.append(.assistant(roundText)) } guard !collectedToolCalls.isEmpty else { break } @@ -1365,7 +1377,11 @@ import Foundation } previousToolCallSignature = signature - let resolution = try await resolveToolCalls(collectedToolCalls, session: session) + let resolution = try await resolveToolCalls( + collectedToolCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -1381,7 +1397,7 @@ import Foundation ) for invocation in invocations { accumulatedEntries.append(.toolOutput(invocation.output)) - chat.append(.tool(toolOutputToJSON(invocation.output))) + pendingChat.append(.tool(toolOutputToJSON(invocation.output))) } yieldSnapshot() } @@ -1438,11 +1454,12 @@ import Foundation let loaded = try await loadContext(modelId: modelId, hub: hub, directory: directory) defer { withExtendedLifetime(loaded) {} } let context = loaded.context - guard let instructions = session.instructions?.description, !instructions.isEmpty else { + let requestContext = session.resolvedRequestContext() + guard let instructions = requestContext.instructions?.description, !instructions.isEmpty else { return } - let toolSpecs = mlxToolSpecs(for: session) + let toolSpecs = mlxToolSpecs(for: requestContext.tools) let params = toGenerateParameters(.init()) let newCache = context.model.newCache(parameters: params) @@ -1533,27 +1550,27 @@ import Foundation // MARK: - Transcript Conversion private func convertTranscriptToMLXChat( - session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, fallbackPrompt: String ) -> [MLXLMCommon.Chat.Message] { var chat: [MLXLMCommon.Chat.Message] = [] // Check if instructions are already in transcript - let hasInstructionsInTranscript = session.transcript.contains { + let hasInstructionsInTranscript = requestContext.transcript.contains { if case .instructions = $0 { return true } return false } // Add instructions from session if present and not in transcript if !hasInstructionsInTranscript, - let instructions = session.instructions?.description, + let instructions = requestContext.instructions?.description, !instructions.isEmpty { chat.append(.init(role: .system, content: instructions)) } // Convert each transcript entry - for entry in session.transcript { + for entry in requestContext.transcript { switch entry { case .instructions(let instr): chat.append(makeMLXChatMessage(from: instr.segments, role: .system)) @@ -1742,12 +1759,13 @@ import Foundation private func resolveToolCalls( _ toolCalls: [MLXLMCommon.ToolCall], + tools: [any Tool], session: LanguageModelSession ) async throws -> ToolResolutionOutcome { if toolCalls.isEmpty { return .invocations([]) } var toolsByName: [String: any Tool] = [:] - for tool in session.tools { + for tool in tools { if toolsByName[tool.name] == nil { toolsByName[tool.name] = tool } @@ -1909,7 +1927,7 @@ import Foundation private func generateStructuredJSON( context: ModelContext, tokenCache: StructuredGenerationTokenCache, - session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, prompt: Prompt, schema: GenerationSchema, options: GenerationOptions, @@ -1918,7 +1936,10 @@ import Foundation let maxTokens = options.maximumResponseTokens ?? 512 let generateParameters = toStructuredGenerateParameters(options) - let baseChat = convertTranscriptToMLXChat(session: session, fallbackPrompt: prompt.description) + let baseChat = convertTranscriptToMLXChat( + requestContext: requestContext, + fallbackPrompt: prompt.description + ) let schemaPrompt = includeSchemaInPrompt ? schemaPrompt(for: schema) : nil let chat = normalizeChatForStructuredGeneration(baseChat, schemaPrompt: schemaPrompt) diff --git a/Sources/AnyLanguageModel/Models/OllamaLanguageModel.swift b/Sources/AnyLanguageModel/Models/OllamaLanguageModel.swift index fcbb7f11..e4d0b12d 100644 --- a/Sources/AnyLanguageModel/Models/OllamaLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/OllamaLanguageModel.swift @@ -116,19 +116,7 @@ public struct OllamaLanguageModel: LanguageModel { includeSchemaInPrompt: Bool, options: GenerationOptions ) async throws -> LanguageModelSession.Response where Content: Generable { - let userSegments = extractPromptSegments(from: session, fallbackText: prompt.description) - let (ollamaText, ollamaImages) = convertSegmentsToOllama(userSegments) - let messages = [ - OllamaMessage( - role: .user, - content: ollamaText, - images: ollamaImages.isEmpty ? nil : ollamaImages - ) - ] let ollamaOptions = convertOptions(options) - let ollamaTools = try session.tools.map { tool in - try convertToolToOllamaFormat(tool) - } let ollamaFormat: JSONValue? if type == String.self { ollamaFormat = nil @@ -137,6 +125,12 @@ public struct OllamaLanguageModel: LanguageModel { ollamaFormat = try JSONValue(schema) } + let requestContext = session.resolvedRequestContext() + let ollamaTools = try requestContext.tools.map(convertToolToOllamaFormat) + var messages = try requestContext.transcript.toOllamaMessages() + if messages.isEmpty { + messages.append(.init(role: .user, content: prompt.description)) + } let params = try createChatParams( model: model, messages: messages, @@ -160,7 +154,11 @@ public struct OllamaLanguageModel: LanguageModel { let usage = chatResponse.reportedUsage?.value ?? .zero if let toolCalls = chatResponse.message.toolCalls, !toolCalls.isEmpty { - let resolution = try await resolveToolCalls(toolCalls, session: session) + let resolution = try await resolveToolCalls( + toolCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -249,16 +247,19 @@ public struct OllamaLanguageModel: LanguageModel { continuation in let task = Task { do { - let tools = try session.tools.map { try convertToolToOllamaFormat($0) } let format = type == String.self ? nil : try JSONValue(convertSchemaToOllamaFormat(schema)) - var messages = try session.transcript.toOllamaMessages() - if messages.isEmpty { - messages.append(.init(role: .user, content: prompt.description)) - } + var inFlightMessages: [OllamaMessage] = [] var state = StreamingResponseState() var toolRounds = ToolRoundLimit(provider: "Ollama") while true { try Task.checkCancellation() + let requestContext = session.resolvedRequestContext() + let tools = try requestContext.tools.map(convertToolToOllamaFormat) + var messages = try requestContext.transcript.toOllamaMessages() + if messages.isEmpty { + messages.append(.init(role: .user, content: prompt.description)) + } + messages.append(contentsOf: inFlightMessages) let params = try createChatParams( model: model, messages: messages, @@ -288,14 +289,18 @@ public struct OllamaLanguageModel: LanguageModel { guard !toolCalls.isEmpty else { break } try Task.checkCancellation() try toolRounds.record(toolCalls.map(\.roundCall)) - switch try await resolveToolCalls(toolCalls, session: session) { + switch try await resolveToolCalls( + toolCalls, + tools: requestContext.tools, + session: session + ) { case .stop(let calls): state.entries.append(.toolCalls(Transcript.ToolCalls(calls))) continuation.yield(try state.stoppedSnapshot()) continuation.finish() return case .invocations(let invocations): - messages.append( + inFlightMessages.append( .init( role: .assistant, content: state.text, @@ -306,7 +311,7 @@ public struct OllamaLanguageModel: LanguageModel { for invocation in invocations { state.entries.append(.toolOutput(invocation.output)) let (text, images) = convertSegmentsToOllama(invocation.output.segments) - messages.append( + inFlightMessages.append( .init( role: .tool, content: text, @@ -344,6 +349,7 @@ private enum ToolResolutionOutcome { private func resolveToolCalls( _ toolCalls: [OllamaToolCall], + tools: [any Tool], session: LanguageModelSession ) async throws -> ToolResolutionOutcome { if toolCalls.isEmpty { @@ -351,7 +357,7 @@ private func resolveToolCalls( } var toolsByName: [String: any Tool] = [:] - for tool in session.tools { + for tool in tools { if toolsByName[tool.name] == nil { toolsByName[tool.name] = tool } @@ -666,15 +672,6 @@ private func convertSegmentsToOllama(_ segments: [Transcript.Segment]) -> (Strin return (textParts.joined(separator: "\n"), images) } -private func extractPromptSegments(from session: LanguageModelSession, fallbackText: String) -> [Transcript.Segment] { - for entry in session.transcript.reversed() { - if case .prompt(let p) = entry { - return p.segments - } - } - return [.text(.init(content: fallbackText))] -} - private struct ChatResponse: Decodable, Sendable { let model: String let createdAt: Date diff --git a/Sources/AnyLanguageModel/Models/OpenAILanguageModel.swift b/Sources/AnyLanguageModel/Models/OpenAILanguageModel.swift index ac297f9f..f92701b7 100644 --- a/Sources/AnyLanguageModel/Models/OpenAILanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/OpenAILanguageModel.swift @@ -465,22 +465,9 @@ public struct OpenAILanguageModel: LanguageModel { includeSchemaInPrompt: Bool, options: GenerationOptions ) async throws -> LanguageModelSession.Response where Content: Generable { - // Convert tools if any are available in the session - let openAITools: [OpenAITool]? = { - guard !session.tools.isEmpty else { return nil } - var converted: [OpenAITool] = [] - converted.reserveCapacity(session.tools.count) - for tool in session.tools { - converted.append(convertToolToOpenAIFormat(tool)) - } - return converted - }() - switch apiVariant { case .chatCompletions: return try await respondWithChatCompletions( - messages: session.transcript.toOpenAIMessages(), - tools: openAITools, generating: type, schema: schema, options: options, @@ -488,8 +475,6 @@ public struct OpenAILanguageModel: LanguageModel { ) case .responses: return try await respondWithResponses( - messages: session.transcript.toOpenAIMessages(), - tools: openAITools, generating: type, schema: schema, options: options, @@ -499,8 +484,6 @@ public struct OpenAILanguageModel: LanguageModel { } private func respondWithChatCompletions( - messages: [OpenAIMessage], - tools: [OpenAITool]?, generating type: Content.Type, schema: GenerationSchema, options: GenerationOptions, @@ -512,11 +495,16 @@ public struct OpenAILanguageModel: LanguageModel { var text = "" // The text of earlier tool rounds, which string responses include. var earlierText = "" - var messages = messages + var inFlightMessages: [OpenAIMessage] = [] var toolRounds = ToolRoundLimit(provider: "OpenAI") // Loop until no more tool calls while true { + let requestContext = session.resolvedRequestContext() + let tools = + requestContext.tools.isEmpty + ? nil : requestContext.tools.map(convertToolToOpenAIFormat) + let messages = requestContext.transcript.toOpenAIMessages() + inFlightMessages let params = try ChatCompletions.createRequestBody( model: model, messages: messages, @@ -559,10 +547,14 @@ public struct OpenAILanguageModel: LanguageModel { let toolCallMessage = choice.message if let toolCalls = toolCallMessage.toolCalls, !toolCalls.isEmpty { if let value = try? JSONValue(toolCallMessage) { - messages.append(OpenAIMessage(role: .raw(rawContent: value), content: .text(""))) + inFlightMessages.append(OpenAIMessage(role: .raw(rawContent: value), content: .text(""))) } try toolRounds.record(toolCalls.map(\.roundCall)) - let resolution = try await resolveToolCalls(toolCalls, session: session) + let resolution = try await resolveToolCalls( + toolCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -581,7 +573,7 @@ public struct OpenAILanguageModel: LanguageModel { for invocation in invocations { let output = invocation.output entries.append(.toolOutput(output)) - messages.append( + inFlightMessages.append( OpenAIMessage( role: .tool(id: invocation.call.id), content: .text(convertSegmentsToToolContentString(output.segments)) @@ -618,8 +610,6 @@ public struct OpenAILanguageModel: LanguageModel { } private func respondWithResponses( - messages: [OpenAIMessage], - tools: [OpenAITool]?, generating type: Content.Type, schema: GenerationSchema, options: GenerationOptions, @@ -631,13 +621,18 @@ public struct OpenAILanguageModel: LanguageModel { // The text of earlier tool rounds, which string responses include. var earlierText = "" var lastOutput: [JSONValue]? - var messages = messages + var inFlightMessages: [OpenAIMessage] = [] let url = baseURL.appendingPathComponent("responses") var toolRounds = ToolRoundLimit(provider: "OpenAI") // Loop until no more tool calls while true { + let requestContext = session.resolvedRequestContext() + let tools = + requestContext.tools.isEmpty + ? nil : requestContext.tools.map(convertToolToOpenAIFormat) + let messages = requestContext.transcript.toOpenAIMessages() + inFlightMessages let params = try Responses.createRequestBody( model: model, messages: messages, @@ -666,11 +661,15 @@ public struct OpenAILanguageModel: LanguageModel { if !toolCalls.isEmpty { if let output = resp.output { for msg in output { - messages.append(OpenAIMessage(role: .raw(rawContent: msg), content: .text(""))) + inFlightMessages.append(OpenAIMessage(role: .raw(rawContent: msg), content: .text(""))) } } try toolRounds.record(toolCalls.map(\.roundCall)) - let resolution = try await resolveToolCalls(toolCalls, session: session) + let resolution = try await resolveToolCalls( + toolCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -690,7 +689,7 @@ public struct OpenAILanguageModel: LanguageModel { for invocation in invocations { let output = invocation.output entries.append(.toolOutput(output)) - messages.append( + inFlightMessages.append( OpenAIMessage( role: .tool(id: invocation.call.id), content: .text(convertSegmentsToToolContentString(output.segments)) @@ -774,16 +773,20 @@ public struct OpenAILanguageModel: LanguageModel { includeSchemaInPrompt: Bool, options: GenerationOptions ) -> sending LanguageModelSession.ResponseStream where Content: Generable { - let tools = session.tools.isEmpty ? nil : session.tools.map(convertToolToOpenAIFormat) let stream = AsyncThrowingStream.Snapshot, any Error> { continuation in let task = Task { do { - var messages = session.transcript.toOpenAIMessages() + var inFlightMessages: [OpenAIMessage] = [] var state = StreamingResponseState() var toolRounds = ToolRoundLimit(provider: "OpenAI") while true { try Task.checkCancellation() + let requestContext = session.resolvedRequestContext() + let tools = + requestContext.tools.isEmpty + ? nil : requestContext.tools.map(convertToolToOpenAIFormat) + let messages = requestContext.transcript.toOpenAIMessages() + inFlightMessages let params: JSONValue let path: String switch apiVariant { @@ -832,7 +835,9 @@ public struct OpenAILanguageModel: LanguageModel { toolCalls = extractToolCallsFromOutput(response?.output) if !toolCalls.isEmpty, let output = response?.output { for item in output { - messages.append(.init(role: .raw(rawContent: item), content: .text(""))) + inFlightMessages.append( + .init(role: .raw(rawContent: item), content: .text("")) + ) } } if let snapshot = state.snapshot() { continuation.yield(snapshot) } @@ -871,14 +876,20 @@ public struct OpenAILanguageModel: LanguageModel { "role": .string("assistant"), "content": .string(state.text), "tool_calls": try JSONValue(toolCalls), ]) - messages.append(.init(role: .raw(rawContent: message), content: .text(""))) + inFlightMessages.append( + .init(role: .raw(rawContent: message), content: .text("")) + ) } } guard !toolCalls.isEmpty else { break } try Task.checkCancellation() try toolRounds.record(toolCalls.map(\.roundCall)) - switch try await resolveToolCalls(toolCalls, session: session) { + switch try await resolveToolCalls( + toolCalls, + tools: requestContext.tools, + session: session + ) { case .stop(let calls): state.entries.append(.toolCalls(Transcript.ToolCalls(calls))) continuation.yield(try state.stoppedSnapshot()) @@ -892,7 +903,7 @@ public struct OpenAILanguageModel: LanguageModel { state.entries.append(.toolCalls(Transcript.ToolCalls(invocations.map(\.call)))) for invocation in invocations { state.entries.append(.toolOutput(invocation.output)) - messages.append( + inFlightMessages.append( .init( role: .tool(id: invocation.call.id), content: .text(convertSegmentsToToolContentString(invocation.output.segments)) @@ -1719,12 +1730,13 @@ private enum OpenAIToolResolutionOutcome { private func resolveToolCalls( _ toolCalls: [OpenAIToolCall], + tools: [any Tool], session: LanguageModelSession ) async throws -> OpenAIToolResolutionOutcome { if toolCalls.isEmpty { return .invocations([]) } var toolsByName: [String: any Tool] = [:] - for tool in session.tools { + for tool in tools { if toolsByName[tool.name] == nil { toolsByName[tool.name] = tool } diff --git a/Sources/AnyLanguageModel/Models/OpenResponsesLanguageModel.swift b/Sources/AnyLanguageModel/Models/OpenResponsesLanguageModel.swift index 82dfb46a..00cacc93 100644 --- a/Sources/AnyLanguageModel/Models/OpenResponsesLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/OpenResponsesLanguageModel.swift @@ -432,11 +432,7 @@ public struct OpenResponsesLanguageModel: LanguageModel { includeSchemaInPrompt: Bool, options: GenerationOptions ) async throws -> LanguageModelSession.Response where Content: Generable { - let tools: [OpenResponsesTool]? = - session.tools.isEmpty ? nil : session.tools.map { convertToolToOpenResponsesFormat($0) } return try await respondWithOpenResponses( - messages: session.transcript.toOpenResponsesMessages(), - tools: tools, generating: type, schema: schema, options: options, @@ -486,18 +482,21 @@ public struct OpenResponsesLanguageModel: LanguageModel { includeSchemaInPrompt: Bool, options: GenerationOptions ) -> sending LanguageModelSession.ResponseStream where Content: Generable { - let tools: [OpenResponsesTool]? = - session.tools.isEmpty ? nil : session.tools.map { convertToolToOpenResponsesFormat($0) } let url = baseURL.appendingPathComponent("responses") let stream = AsyncThrowingStream.Snapshot, any Error> { continuation in let task = Task { do { - var messages = session.transcript.toOpenResponsesMessages() + var inFlightMessages: [OpenResponsesMessage] = [] var state = StreamingResponseState() var toolRounds = ToolRoundLimit(provider: "Open Responses") while true { try Task.checkCancellation() + let requestContext = session.resolvedRequestContext() + let tools: [OpenResponsesTool]? = + requestContext.tools.isEmpty + ? nil : requestContext.tools.map(convertToolToOpenResponsesFormat) + let messages = requestContext.transcript.toOpenResponsesMessages() + inFlightMessages let params = try OpenResponsesAPI.createRequestBody( model: model, messages: messages, @@ -528,7 +527,9 @@ public struct OpenResponsesLanguageModel: LanguageModel { toolCalls = extractToolCallsFromOutput(response?.output) if !toolCalls.isEmpty, let output = response?.output { for item in output { - messages.append(.init(role: .raw(rawContent: item), content: .text(""))) + inFlightMessages.append( + .init(role: .raw(rawContent: item), content: .text("")) + ) } } if let snapshot = state.snapshot() { continuation.yield(snapshot) } @@ -545,7 +546,11 @@ public struct OpenResponsesLanguageModel: LanguageModel { guard !toolCalls.isEmpty else { break } try Task.checkCancellation() try toolRounds.record(toolCalls.map(\.roundCall)) - switch try await resolveToolCalls(toolCalls, session: session) { + switch try await resolveToolCalls( + toolCalls, + tools: requestContext.tools, + session: session + ) { case .stop(let calls): state.entries.append(.toolCalls(Transcript.ToolCalls(calls))) continuation.yield(try state.stoppedSnapshot()) @@ -555,7 +560,7 @@ public struct OpenResponsesLanguageModel: LanguageModel { state.entries.append(.toolCalls(Transcript.ToolCalls(invocations.map(\.call)))) for invocation in invocations { state.entries.append(.toolOutput(invocation.output)) - messages.append( + inFlightMessages.append( .init( role: .tool(id: invocation.call.id), content: .text( @@ -580,8 +585,6 @@ public struct OpenResponsesLanguageModel: LanguageModel { /// Sends a non-streaming request to the Open Responses API and returns the parsed response. private func respondWithOpenResponses( - messages: [OpenResponsesMessage], - tools: [OpenResponsesTool]?, generating type: Content.Type, schema: GenerationSchema, options: GenerationOptions, @@ -593,11 +596,16 @@ public struct OpenResponsesLanguageModel: LanguageModel { // The text of earlier tool rounds, which string responses include. var earlierText = "" var lastOutput: [JSONValue]? - var messages = messages + var inFlightMessages: [OpenResponsesMessage] = [] let url = baseURL.appendingPathComponent("responses") var toolRounds = ToolRoundLimit(provider: "Open Responses") while true { + let requestContext = session.resolvedRequestContext() + let tools: [OpenResponsesTool]? = + requestContext.tools.isEmpty + ? nil : requestContext.tools.map(convertToolToOpenResponsesFormat) + let messages = requestContext.transcript.toOpenResponsesMessages() + inFlightMessages let params = try OpenResponsesAPI.createRequestBody( model: model, messages: messages, @@ -622,11 +630,17 @@ public struct OpenResponsesLanguageModel: LanguageModel { if !toolCalls.isEmpty { if let output = resp.output { for item in output { - messages.append(OpenResponsesMessage(role: .raw(rawContent: item), content: .text(""))) + inFlightMessages.append( + OpenResponsesMessage(role: .raw(rawContent: item), content: .text("")) + ) } } try toolRounds.record(toolCalls.map(\.roundCall)) - let resolution = try await resolveToolCalls(toolCalls, session: session) + let resolution = try await resolveToolCalls( + toolCalls, + tools: requestContext.tools, + session: session + ) switch resolution { case .stop(let calls): if !calls.isEmpty { @@ -644,7 +658,7 @@ public struct OpenResponsesLanguageModel: LanguageModel { entries.append(.toolCalls(Transcript.ToolCalls(invocations.map { $0.call }))) for inv in invocations { entries.append(.toolOutput(inv.output)) - messages.append( + inFlightMessages.append( OpenResponsesMessage( role: .tool(id: inv.call.id), content: .text(openResponsesConvertSegmentsToToolContentString(inv.output.segments)) @@ -1154,11 +1168,12 @@ private enum OpenResponsesToolResolutionOutcome: Sendable { private func resolveToolCalls( _ toolCalls: [OpenResponsesToolCall], + tools: [any Tool], session: LanguageModelSession ) async throws -> OpenResponsesToolResolutionOutcome { if toolCalls.isEmpty { return .invocations([]) } var byName: [String: any Tool] = [:] - for t in session.tools { if byName[t.name] == nil { byName[t.name] = t } } + for t in tools { if byName[t.name] == nil { byName[t.name] = t } } var transcriptCalls: [Transcript.ToolCall] = [] for c in toolCalls { let args = (c.arguments.flatMap { try? GeneratedContent(json: $0) } ?? GeneratedContent(properties: [:])) diff --git a/Sources/AnyLanguageModel/Models/PrivateCloudComputeLanguageModel.swift b/Sources/AnyLanguageModel/Models/PrivateCloudComputeLanguageModel.swift index 7f8f1aa9..3681aa0b 100644 --- a/Sources/AnyLanguageModel/Models/PrivateCloudComputeLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/PrivateCloudComputeLanguageModel.swift @@ -126,10 +126,11 @@ issues: [LanguageModelFeedback.Issue], desiredOutput: Transcript.Entry? ) -> Data { + let requestContext = session.resolvedRequestContext() let fmSession = FoundationModels.LanguageModelSession( model: pccModel, - tools: session.tools.toFoundationModels(), - instructions: session.instructions?.toFoundationModels() + tools: requestContext.tools.toFoundationModels(), + instructions: requestContext.instructions?.toFoundationModels() ) return fmSession.logFeedbackAttachment( sentiment: sentiment?.toFoundationModels(), diff --git a/Sources/AnyLanguageModel/Models/SystemLanguageModel.swift b/Sources/AnyLanguageModel/Models/SystemLanguageModel.swift index d4f10119..949b0367 100644 --- a/Sources/AnyLanguageModel/Models/SystemLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/SystemLanguageModel.swift @@ -128,19 +128,8 @@ let fmPrompt = prompt.toFoundationModels() let fmOptions = options.toFoundationModels() - let fmSession = FoundationModels.LanguageModelSession( - model: systemModel, - tools: session.tools.toFoundationModels(), - transcript: fmTranscriptDroppingDuplicatePrompt(session.transcript, prompt: prompt).toFoundationModels( - instructions: session.instructions, - toolDefinitions: session.tools - .filter(\.includesSchemaInInstructions) - .map { Transcript.ToolDefinition(tool: $0) } - ) - ) - return try await fmRespond( - makeSession: { fmSession }, + makeSession: { try self.makeSession(for: session, prompt: prompt) }, fmPrompt: fmPrompt, fmOptions: fmOptions, type: type, @@ -194,19 +183,8 @@ let fmPrompt = prompt.toFoundationModels() let fmOptions = options.toFoundationModels() - let fmSession = FoundationModels.LanguageModelSession( - model: systemModel, - tools: session.tools.toFoundationModels(), - transcript: fmTranscriptDroppingDuplicatePrompt(session.transcript, prompt: prompt).toFoundationModels( - instructions: session.instructions, - toolDefinitions: session.tools - .filter(\.includesSchemaInInstructions) - .map { Transcript.ToolDefinition(tool: $0) } - ) - ) - return fmStreamResponse( - makeSession: { fmSession }, + makeSession: { try self.makeSession(for: session, prompt: prompt) }, fmPrompt: fmPrompt, fmOptions: fmOptions, type: type, @@ -221,10 +199,11 @@ issues: [LanguageModelFeedback.Issue], desiredOutput: Transcript.Entry? ) -> Data { + let requestContext = session.resolvedRequestContext() let fmSession = FoundationModels.LanguageModelSession( model: systemModel, - tools: session.tools.toFoundationModels(), - instructions: session.instructions?.toFoundationModels() + tools: requestContext.tools.toFoundationModels(), + instructions: requestContext.instructions?.toFoundationModels() ) let fmSentiment = sentiment?.toFoundationModels() @@ -238,6 +217,39 @@ ) } + private func makeSession( + for session: LanguageModelSession, + prompt: Prompt + ) throws -> FoundationModels.LanguageModelSession { + #if compiler(>=6.4) && !os(tvOS) + if #available(macOS 27.0, iOS 27.0, visionOS 27.0, watchOS 27.0, *) { + return makeFoundationModelsSession( + model: systemModel, + session: session, + prompt: prompt + ) + } + #endif + + guard !session.usesDynamicInstructions else { + throw foundationModelsDynamicInstructionsUnavailableError() + } + let requestContext = session.resolvedRequestContext() + return FoundationModels.LanguageModelSession( + model: systemModel, + tools: requestContext.tools.toFoundationModels(), + transcript: fmTranscriptDroppingDuplicatePrompt( + requestContext.transcript, + prompt: prompt + ).toFoundationModels( + instructions: requestContext.instructions, + toolDefinitions: requestContext.tools + .filter(\.includesSchemaInInstructions) + .map { Transcript.ToolDefinition(tool: $0) } + ) + ) + } + } // MARK: - Helpers @@ -257,6 +269,69 @@ return Transcript(entries: transcript.dropLast()) } + func foundationModelsDynamicInstructionsUnavailableError() -> LanguageModelSession.GenerationError { + .decodingFailure( + .init( + debugDescription: + "Dynamic instructions require the Foundation Models 27 runtime for native request and tool-continuation semantics." + ) + ) + } + + #if compiler(>=6.4) && !os(tvOS) + @available(macOS 27.0, iOS 27.0, visionOS 27.0, watchOS 27.0, *) + struct FoundationModelsDynamicInstructionsAdapter: FoundationModels.DynamicInstructions { + let session: LanguageModelSession + + var body: some FoundationModels.DynamicInstructions { + let requestContext = session.resolvedRequestContext() + if let instructions = requestContext.instructions { + instructions.toFoundationModels() + } + requestContext.tools.toFoundationModels() + } + } + + @available(macOS 27.0, iOS 27.0, visionOS 27.0, watchOS 27.0, *) + func makeFoundationModelsSession( + model: Model, + session: LanguageModelSession, + prompt: Prompt + ) -> FoundationModels.LanguageModelSession { + if session.usesDynamicInstructions { + let history = fmTranscriptDroppingDuplicatePrompt( + Transcript( + entries: session.transcript.filter { entry in + if case .instructions = entry { return false } + return true + } + ), + prompt: prompt + ).toFoundationModels(instructions: nil, toolDefinitions: []) + return FoundationModels.LanguageModelSession( + model: model, + dynamicInstructions: FoundationModelsDynamicInstructionsAdapter(session: session), + history: history + ) + } + + let requestContext = session.resolvedRequestContext() + return FoundationModels.LanguageModelSession( + model: model, + tools: requestContext.tools.toFoundationModels(), + transcript: fmTranscriptDroppingDuplicatePrompt( + requestContext.transcript, + prompt: prompt + ).toFoundationModels( + instructions: requestContext.instructions, + toolDefinitions: requestContext.tools + .filter(\.includesSchemaInInstructions) + .map { Transcript.ToolDefinition(tool: $0) } + ) + ) + } + #endif + @available(macOS 26.0, iOS 26.0, watchOS 27.0, tvOS 26.0, visionOS 26.0, *) extension Prompt { func toFoundationModels() -> FoundationModels.Prompt { diff --git a/Tests/AnyLanguageModelTests/APICompatibilityAnyLanguageModelTests.swift b/Tests/AnyLanguageModelTests/APICompatibilityAnyLanguageModelTests.swift index 55106cdd..18bf1cba 100644 --- a/Tests/AnyLanguageModelTests/APICompatibilityAnyLanguageModelTests.swift +++ b/Tests/AnyLanguageModelTests/APICompatibilityAnyLanguageModelTests.swift @@ -10,6 +10,18 @@ import Testing return false }() + @available(macOS 27.0, iOS 27.0, visionOS 27.0, *) + private struct CompatibilityDynamicInstructions: DynamicInstructions { + let includeDetail: Bool + + var body: some DynamicInstructions { + Instructions("You are a helpful assistant.") + if includeDetail { + Instructions("Include useful detail.") + } + } + } + @available(macOS 26.0, iOS 26.0, tvOS 26.0, visionOS 26.0, *) @Test("AnyLanguageModel Drop-In Compatibility", .enabled(if: isSystemLanguageModelAvailable)) func anyLanguageModelCompatibility() async throws { @@ -19,6 +31,14 @@ import Testing instructions: Instructions("You are a helpful assistant.") ) + if #available(macOS 27.0, iOS 27.0, visionOS 27.0, *) { + _ = LanguageModelSession( + model: model, + dynamicInstructions: CompatibilityDynamicInstructions(includeDetail: true), + history: session.transcript + ) + } + let options = GenerationOptions(temperature: 0.7) let response = try await session.respond(options: options) { Prompt("Say 'Hello'") diff --git a/Tests/AnyLanguageModelTests/APICompatibilityFoundationModelsTests.swift b/Tests/AnyLanguageModelTests/APICompatibilityFoundationModelsTests.swift index 60943da5..8142648e 100644 --- a/Tests/AnyLanguageModelTests/APICompatibilityFoundationModelsTests.swift +++ b/Tests/AnyLanguageModelTests/APICompatibilityFoundationModelsTests.swift @@ -10,6 +10,22 @@ import Testing return false }() + #if compiler(>=6.4) + #if os(macOS) || os(iOS) || os(visionOS) + @available(macOS 27.0, iOS 27.0, visionOS 27.0, *) + private struct CompatibilityDynamicInstructions: DynamicInstructions { + let includeDetail: Bool + + var body: some DynamicInstructions { + Instructions("You are a helpful assistant.") + if includeDetail { + Instructions("Include useful detail.") + } + } + } + #endif + #endif + @available(macOS 26.0, iOS 26.0, tvOS 26.0, visionOS 26.0, *) @Test( "FoundationModels Drop-In Compatibility", @@ -22,6 +38,18 @@ import Testing instructions: Instructions("You are a helpful assistant.") ) + #if compiler(>=6.4) + #if os(macOS) || os(iOS) || os(visionOS) + if #available(macOS 27.0, iOS 27.0, visionOS 27.0, *) { + _ = LanguageModelSession( + model: model, + dynamicInstructions: CompatibilityDynamicInstructions(includeDetail: true), + history: session.transcript + ) + } + #endif + #endif + let options = GenerationOptions(temperature: 0.7) let response = try await session.respond(options: options) { Prompt("Say 'Hello'") diff --git a/Tests/AnyLanguageModelTests/DynamicInstructionsTests.swift b/Tests/AnyLanguageModelTests/DynamicInstructionsTests.swift new file mode 100644 index 00000000..e45fe6ca --- /dev/null +++ b/Tests/AnyLanguageModelTests/DynamicInstructionsTests.swift @@ -0,0 +1,571 @@ +import Foundation +import Testing + +@testable import AnyLanguageModel + +@Suite("Dynamic instructions") +struct DynamicInstructionsTests { + @Test func bodyReevaluatesForEveryNonstreamingRequest() async throws { + let state = DynamicFixtureState() + let model = DynamicContextModel(state: state, continuesAfterTool: false) + let session = LanguageModelSession( + model: model, + dynamicInstructions: FixtureDynamicInstructions(state: state) + ) + + #expect(state.evaluationCount == 0) + _ = try await session.respond(to: "First") + state.select(.b) + _ = try await session.respond(to: "Second") + + #expect( + model.snapshots.withLock { $0 } == [ + .init(instructions: "Instructions A", tools: ["tool-a"]), + .init(instructions: "Instructions B", tools: ["tool-b"]), + ] + ) + #expect(state.evaluationCount == 2) + } + + @Test func bodyReevaluatesForEveryStreamingRequest() async throws { + let state = DynamicFixtureState() + let model = DynamicContextModel(state: state, continuesAfterTool: false) + let session = LanguageModelSession( + model: model, + dynamicInstructions: FixtureDynamicInstructions(state: state) + ) + + #expect(state.evaluationCount == 0) + _ = try await session.streamResponse(to: "First").collect() + state.select(.b) + _ = try await session.streamResponse(to: "Second").collect() + + #expect( + model.snapshots.withLock { $0 } == [ + .init(instructions: "Instructions A", tools: ["tool-a"]), + .init(instructions: "Instructions B", tools: ["tool-b"]), + ] + ) + #expect(state.evaluationCount == 2) + } + + @Test(arguments: [false, true]) + func toolContinuationReevaluatesAndExecutesProducingSnapshot(streaming: Bool) async throws { + let state = DynamicFixtureState() + let model = DynamicContextModel(state: state, continuesAfterTool: true) + let session = LanguageModelSession( + model: model, + dynamicInstructions: FixtureDynamicInstructions(state: state) + ) + + if streaming { + _ = try await session.streamResponse(to: "Use a tool").collect() + } else { + _ = try await session.respond(to: "Use a tool") + } + + #expect( + model.snapshots.withLock { $0 } == [ + .init(instructions: "Instructions A", tools: ["tool-a"]), + .init(instructions: "Instructions B", tools: ["tool-b"]), + ] + ) + #expect(state.executedTools == ["tool-a"]) + #expect(state.evaluationCount == 2) + #expect(session.transcript.count == 4) + guard case .prompt = session.transcript[0], + case .toolCalls(let calls) = session.transcript[1], + case .toolOutput(let output) = session.transcript[2], + case .response = session.transcript[3] + else { + Issue.record("Expected prompt, tool call, tool output, and response") + return + } + #expect(calls.first?.id == output.id) + #expect(output.toolName == "tool-a") + } + + @Test func historyDoesNotPersistDynamicInstructionsAndRehydratesWithCurrentState() async throws { + let state = DynamicFixtureState() + let firstModel = DynamicContextModel(state: state, continuesAfterTool: false) + let firstSession = LanguageModelSession( + model: firstModel, + dynamicInstructions: FixtureDynamicInstructions(state: state) + ) + _ = try await firstSession.respond(to: "First") + + #expect( + !firstSession.transcript.contains { + if case .instructions = $0 { true } else { false } + } + ) + + state.select(.b) + let restoredModel = DynamicContextModel(state: state, continuesAfterTool: false) + let restoredSession = LanguageModelSession( + model: restoredModel, + dynamicInstructions: FixtureDynamicInstructions(state: state), + history: firstSession.transcript + ) + _ = try await restoredSession.respond(to: "Second") + + #expect( + restoredModel.snapshots.withLock { $0 } == [ + .init(instructions: "Instructions B", tools: ["tool-b"]) + ] + ) + #expect(restoredSession.transcript.count == 4) + #expect( + !restoredSession.transcript.contains { + if case .instructions = $0 { true } else { false } + } + ) + } + + @Test func builderComposesNestedConditionalEmptyAndToolArrayContent() { + let state = DynamicFixtureState() + let enabled = true + let dynamic = AnyDynamicInstructions(erasing: FixtureComposition(state: state, enabled: enabled)) + let session = LanguageModelSession( + model: DynamicContextModel(state: state, continuesAfterTool: false), + dynamicInstructions: dynamic + ) + + let context = session.resolvedRequestContext() + + #expect(context.instructions?.description == "Outer\nNested\nFor each") + #expect(context.tools.map(\.name) == ["tool-a", "tool-b"]) + #expect(session.transcript.isEmpty) + } + + @Test func staticSessionRequestContextPreservesExistingBehavior() { + let state = DynamicFixtureState() + let tool = FixtureTool(name: "static-tool", state: state) + let session = LanguageModelSession( + model: DynamicContextModel(state: state, continuesAfterTool: false), + tools: [tool], + instructions: "Static" + ) + + let context = session.resolvedRequestContext() + + #expect(context.instructions?.description == "Static") + #expect(context.tools.map(\.name) == ["static-tool"]) + #expect(context.transcript == session.transcript) + #expect(session.instructions?.description == "Static") + #expect(session.tools.map(\.name) == ["static-tool"]) + } + + @Test func failedResponseDoesNotReplayCompletedDynamicToolSideEffect() async throws { + let state = DynamicFixtureState() + let model = FailingAfterToolModel(state: state) + let session = LanguageModelSession( + model: model, + dynamicInstructions: FixtureDynamicInstructions(state: state) + ) + + await #expect(throws: DynamicFixtureError.failed) { + _ = try await session.respond(to: "Fail after tool") + } + _ = try await session.respond(to: "Retry") + + #expect(state.executedTools == ["tool-a"]) + #expect( + model.snapshots.withLock { $0 } == [ + .init(instructions: "Instructions A", tools: ["tool-a"]), + .init(instructions: "Instructions B", tools: ["tool-b"]), + ] + ) + } + + @Test func cancelledResponseDoesNotReplayCompletedDynamicToolSideEffect() async throws { + let state = DynamicFixtureState() + let control = DynamicCancellationControl() + let model = CancellingAfterToolModel(state: state, control: control) + let session = LanguageModelSession( + model: model, + dynamicInstructions: FixtureDynamicInstructions(state: state) + ) + var started = control.started.makeAsyncIterator() + + let response = Task { + try await session.respond(to: "Cancel after tool") + } + _ = await started.next() + response.cancel() + + await #expect(throws: CancellationError.self) { + _ = try await response.value + } + _ = try await session.respond(to: "Retry") + + #expect(state.executedTools == ["tool-a"]) + #expect( + model.snapshots.withLock { $0 } == [ + .init(instructions: "Instructions A", tools: ["tool-a"]), + .init(instructions: "Instructions B", tools: ["tool-b"]), + ] + ) + } +} + +private struct FixtureDynamicInstructions: DynamicInstructions { + let state: DynamicFixtureState + + var body: some DynamicInstructions { + let snapshot = state.snapshotForEvaluation() + Instructions(snapshot.instructions) + [snapshot.tool] + } +} + +private struct FixtureComposition: DynamicInstructions { + let state: DynamicFixtureState + let enabled: Bool + + var body: some DynamicInstructions { + Instructions("Outer") + if enabled { + NestedFixtureInstructions(state: state) + } + EmptyDynamicInstructions() + ForEach([FixtureInstruction(id: 1, text: "For each")]) { item in + Instructions(item.text) + } + } +} + +private struct FixtureInstruction: Identifiable { + let id: Int + let text: String +} + +private struct NestedFixtureInstructions: DynamicInstructions { + let state: DynamicFixtureState + + var body: some DynamicInstructions { + Instructions("Nested") + [ + FixtureTool(name: "tool-a", state: state), + FixtureTool(name: "tool-b", state: state), + ] as [any Tool] + } +} + +private final class DynamicFixtureState: @unchecked Sendable { + enum Selection: Sendable { + case a + case b + } + + struct Storage: Sendable { + var selection = Selection.a + var evaluationCount = 0 + var executedTools: [String] = [] + } + + struct Snapshot: Sendable { + let instructions: String + let tool: any Tool + } + + private let storage = Locked(Storage()) + + var evaluationCount: Int { + storage.withLock { $0.evaluationCount } + } + + var executedTools: [String] { + storage.withLock { $0.executedTools } + } + + func select(_ selection: Selection) { + storage.withLock { $0.selection = selection } + } + + func snapshotForEvaluation() -> Snapshot { + storage.withLock { storage in + storage.evaluationCount += 1 + switch storage.selection { + case .a: + return Snapshot( + instructions: "Instructions A", + tool: FixtureTool(name: "tool-a", state: self) + ) + case .b: + return Snapshot( + instructions: "Instructions B", + tool: FixtureTool(name: "tool-b", state: self) + ) + } + } + } + + func recordExecution(_ name: String) { + storage.withLock { $0.executedTools.append(name) } + } +} + +private struct FixtureTool: Tool { + let name: String + let description = "Records which request-scoped tool instance executed" + let state: DynamicFixtureState + + typealias Arguments = GeneratedContent + + var parameters: GenerationSchema { + GeneratedContent.generationSchema + } + + func call(arguments: GeneratedContent) async throws -> String { + state.recordExecution(name) + return name + } +} + +private struct DynamicRequestSnapshot: Sendable, Equatable { + let instructions: String? + let tools: [String] +} + +private struct DynamicContextModel: LanguageModel { + typealias UnavailableReason = Never + + let state: DynamicFixtureState + let continuesAfterTool: Bool + let snapshots = Locked<[DynamicRequestSnapshot]>([]) + + func respond( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) async throws -> LanguageModelSession.Response where Content: Generable { + let result = try await run(session: session, type: type) + return .init( + content: result.content, + rawContent: result.raw, + transcriptEntries: result.entries + ) + } + + func streamResponse( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) -> sending LanguageModelSession.ResponseStream where Content: Generable { + let stream = AsyncThrowingStream.Snapshot, any Error> { + continuation in + Task { + do { + let result = try await run(session: session, type: type) + continuation.yield( + .init( + content: result.content.asPartiallyGenerated(), + rawContent: result.raw, + transcriptEntries: ArraySlice(result.entries) + ) + ) + continuation.finish() + } catch { + continuation.finish(throwing: error) + } + } + } + return .init(stream: stream) + } + + private func run( + session: LanguageModelSession, + type: Content.Type + ) async throws -> (content: Content, raw: GeneratedContent, entries: ArraySlice) { + let first = session.resolvedRequestContext() + record(first) + var entries: [Transcript.Entry] = [] + + if continuesAfterTool { + let tool = try #require(first.tools.first) + state.select(.b) + let call = Transcript.ToolCall( + id: "request-a-call", + toolName: tool.name, + arguments: GeneratedContent(properties: [:]) + ) + let output = Transcript.ToolOutput( + id: call.id, + toolName: call.toolName, + segments: try await tool.makeOutputSegments(from: call.arguments) + ) + entries.append(.toolCalls(.init(id: "request-a-calls", [call]))) + entries.append(.toolOutput(output)) + + let continuation = session.resolvedRequestContext() + record(continuation) + } + + let raw = GeneratedContent("Done") + return (try Content(raw), raw, ArraySlice(entries)) + } + + private func record(_ context: LanguageModelSession.RequestContext) { + snapshots.withLock { + $0.append( + .init( + instructions: context.instructions?.description, + tools: context.tools.map(\.name) + ) + ) + } + } +} + +private enum DynamicFixtureError: Error, Equatable { + case failed +} + +private struct FailingAfterToolModel: LanguageModel { + typealias UnavailableReason = Never + + let state: DynamicFixtureState + let snapshots = Locked<[DynamicRequestSnapshot]>([]) + private let didFail = Locked(false) + + func respond( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) async throws -> LanguageModelSession.Response where Content: Generable { + let context = session.resolvedRequestContext() + snapshots.withLock { + $0.append( + .init( + instructions: context.instructions?.description, + tools: context.tools.map(\.name) + ) + ) + } + + let shouldFail = didFail.withLock { didFail in + defer { didFail = true } + return !didFail + } + if shouldFail { + let tool = try #require(context.tools.first) + _ = try await tool.makeOutputSegments(from: GeneratedContent(properties: [:])) + state.select(.b) + throw DynamicFixtureError.failed + } + + let raw = GeneratedContent("Done") + return .init(content: try Content(raw), rawContent: raw, transcriptEntries: []) + } + + func streamResponse( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) -> sending LanguageModelSession.ResponseStream where Content: Generable { + let stream = AsyncThrowingStream.Snapshot, any Error> { + $0.finish(throwing: DynamicFixtureError.failed) + } + return .init(stream: stream) + } +} + +private final class DynamicCancellationControl: @unchecked Sendable { + let started: AsyncStream + + private let startedContinuation: AsyncStream.Continuation + private let cancellationContinuation = Locked?>(nil) + + init() { + (started, startedContinuation) = AsyncStream.makeStream() + } + + func signalStarted() { + startedContinuation.yield(()) + } + + func waitForCancellation() async throws { + try await withTaskCancellationHandler { + try await withCheckedThrowingContinuation { continuation in + let isAlreadyCancelled = cancellationContinuation.withLock { stored in + guard !Task.isCancelled else { return true } + stored = continuation + return false + } + if isAlreadyCancelled { + continuation.resume(throwing: CancellationError()) + } + } + } onCancel: { + let continuation = cancellationContinuation.withLock { stored in + defer { stored = nil } + return stored + } + continuation?.resume(throwing: CancellationError()) + } + } +} + +private struct CancellingAfterToolModel: LanguageModel { + typealias UnavailableReason = Never + + let state: DynamicFixtureState + let control: DynamicCancellationControl + let snapshots = Locked<[DynamicRequestSnapshot]>([]) + private let didSuspend = Locked(false) + + func respond( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) async throws -> LanguageModelSession.Response where Content: Generable { + let context = session.resolvedRequestContext() + snapshots.withLock { + $0.append( + .init( + instructions: context.instructions?.description, + tools: context.tools.map(\.name) + ) + ) + } + + let shouldSuspend = didSuspend.withLock { didSuspend in + defer { didSuspend = true } + return !didSuspend + } + if shouldSuspend { + let tool = try #require(context.tools.first) + _ = try await tool.makeOutputSegments(from: GeneratedContent(properties: [:])) + state.select(.b) + control.signalStarted() + try await control.waitForCancellation() + } + + let raw = GeneratedContent("Done") + return .init(content: try Content(raw), rawContent: raw, transcriptEntries: []) + } + + func streamResponse( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) -> sending LanguageModelSession.ResponseStream where Content: Generable { + let stream = AsyncThrowingStream.Snapshot, any Error> { + $0.finish(throwing: CancellationError()) + } + return .init(stream: stream) + } +} diff --git a/Tests/AnyLanguageModelTests/Shared/MockLanguageModel.swift b/Tests/AnyLanguageModelTests/Shared/MockLanguageModel.swift index 423b3b5d..7fa23bc1 100644 --- a/Tests/AnyLanguageModelTests/Shared/MockLanguageModel.swift +++ b/Tests/AnyLanguageModelTests/Shared/MockLanguageModel.swift @@ -76,7 +76,10 @@ struct MockLanguageModel: LanguageModel { $0.append(Request(schema: schema, includeSchemaInPrompt: includeSchemaInPrompt, options: options)) } - let promptWithInstructions = Prompt("Instructions: \(session.instructions?.description ?? "N/A")\n\(prompt)") + let requestContext = session.resolvedRequestContext() + let promptWithInstructions = Prompt( + "Instructions: \(requestContext.instructions?.description ?? "N/A")\n\(prompt)" + ) let text = try await responseProvider(promptWithInstructions, options) let rawContent = try type == String.self ? GeneratedContent(text) : GeneratedContent(json: text) @@ -134,7 +137,10 @@ struct MockLanguageModel: LanguageModel { $0.append(Request(schema: schema, includeSchemaInPrompt: includeSchemaInPrompt, options: options)) } - let promptWithInstructions = Prompt("Instructions: \(session.instructions?.description ?? "N/A")\n\(prompt)") + let requestContext = session.resolvedRequestContext() + let promptWithInstructions = Prompt( + "Instructions: \(requestContext.instructions?.description ?? "N/A")\n\(prompt)" + ) let stream = AsyncThrowingStream.Snapshot, any Error> { continuation in