From 15efb6df769215d5c847d6bcc063826a5a323ef6 Mon Sep 17 00:00:00 2001 From: Mattt Zmuda Date: Sat, 3 Oct 2026 05:33:31 -0700 Subject: [PATCH 1/3] Resolve a request context before each model request --- README.md | 4 +- .../LanguageModelSession.swift | 43 +++++++++ .../Models/AnthropicLanguageModel.swift | 37 +++++--- .../Models/CoreMLLanguageModel.swift | 40 ++++---- .../Models/FoundationLanguageModel.swift | 30 ++---- .../Models/GeminiLanguageModel.swift | 58 ++++++++---- .../Models/LlamaLanguageModel.swift | 93 +++++++++++-------- .../Models/MLXLanguageModel.swift | 75 +++++++++------ .../Models/OllamaLanguageModel.swift | 59 ++++++------ .../Models/OpenAILanguageModel.swift | 80 +++++++++------- .../Models/OpenResponsesLanguageModel.swift | 49 ++++++---- .../PrivateCloudComputeLanguageModel.swift | 9 +- .../Models/SystemLanguageModel.swift | 79 ++++++++++------ .../RequestContextTests.swift | 29 ++++++ .../Shared/MockLanguageModel.swift | 10 +- 15 files changed, 434 insertions(+), 261 deletions(-) create mode 100644 Tests/AnyLanguageModelTests/RequestContextTests.swift diff --git a/README.md b/README.md index e2b6f2f0..adfd9206 100644 --- a/README.md +++ b/README.md @@ -576,8 +576,8 @@ say which API they follow. [observing and controlling tool calls](#tool-calling). - `transcriptErrorHandlingPolicy` and `waitForResponseCompletion()`: what a transcript keeps when a request fails or is cancelled. -- `LanguageModelSession.tools` and `instructions`: - the session's tools and instructions, +- `LanguageModelSession.tools`, `instructions`, and `resolvedRequestContext()`: + the session's tools, instructions, and the inputs for each request, for language models defined outside AnyLanguageModel. - `Usage` and the `usage` properties: [token usage](#token-usage), diff --git a/Sources/AnyLanguageModel/LanguageModelSession.swift b/Sources/AnyLanguageModel/LanguageModelSession.swift index de49bdcc..7e2c2c41 100644 --- a/Sources/AnyLanguageModel/LanguageModelSession.swift +++ b/Sources/AnyLanguageModel/LanguageModelSession.swift @@ -111,6 +111,49 @@ public final class LanguageModelSession: @unchecked Sendable { /// with the Foundation Models framework. @ObservationIgnored public var toolExecutionDelegate: (any ToolExecutionDelegate)? + /// The transcript, instructions, and tools for one model request. + /// + /// A language model creates one context immediately before each request it sends, + /// by calling ``LanguageModelSession/resolvedRequestContext()``. + /// If that request produces tool calls, + /// run them with the ``tools`` from the same context, + /// and resolve a new context only before the continuation request. + /// + /// - Note: This API is exclusive to AnyLanguageModel + /// and using it means your code is no longer drop-in compatible + /// with the Foundation Models framework. + /// It's public so that language models outside this module + /// can read the inputs for each request. + public struct RequestContext: Sendable { + /// The transcript to send with the request. + public let transcript: Transcript + + /// The instructions for the request, if any. + public let instructions: Instructions? + + /// The tools that the model can call in response to the request. + public let tools: [any Tool] + + fileprivate init(transcript: Transcript, instructions: Instructions?, tools: [any Tool]) { + self.transcript = transcript + self.instructions = instructions + self.tools = tools + } + } + + /// Returns the transcript, instructions, and tools for the next model request. + /// + /// Calling this method doesn't change the session. + /// + /// - Note: This API is exclusive to AnyLanguageModel + /// and using it means your code is no longer drop-in compatible + /// with the Foundation Models framework. + /// It's public so that language models outside this module + /// can read the inputs for each request. + nonisolated public func resolvedRequestContext() -> RequestContext { + RequestContext(transcript: transcript, instructions: instructions, tools: tools) + } + /// Creates a session with a model, tools, /// and instructions that you build with a result builder. /// diff --git a/Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift b/Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift index e1d788db..6a618a79 100644 --- a/Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift @@ -445,9 +445,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) } @@ -455,7 +456,7 @@ public struct AnthropicLanguageModel: LanguageModel { let params = try createMessageParams( model: model, system: nil, - messages: try session.transcript.toAnthropicMessages(), + messages: try requestContext.transcript.toAnthropicMessages(), tools: anthropicTools.isEmpty ? nil : anthropicTools, responseSchema: responseSchema, options: options @@ -486,7 +487,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 { @@ -584,19 +589,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 = try 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 = try requestContext.transcript.toAnthropicMessages() + inFlightMessages var params = try createMessageParams( model: model, system: nil, @@ -692,14 +696,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 { @@ -713,7 +721,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() @@ -893,12 +901,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 3aa74e2d..56b35462 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 { @@ -410,7 +402,7 @@ } private func generateStructuredJSON( - session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, prompt: Prompt, schema: GenerationSchema, options: GenerationOptions, @@ -420,7 +412,7 @@ var generationConfig = toStructuredGenerationConfig(options) let promptTokens = try structuredPromptTokens( - in: session, + requestContext: requestContext, prompt: prompt, schema: schema, includeSchemaInPrompt: includeSchemaInPrompt @@ -457,20 +449,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 0c19e3d0..06caadc4 100644 --- a/Sources/AnyLanguageModel/Models/GeminiLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/GeminiLanguageModel.swift @@ -230,12 +230,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. @@ -244,8 +242,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, @@ -282,7 +289,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 { @@ -300,12 +311,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) } } @@ -402,16 +413,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, @@ -452,7 +469,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()) @@ -462,11 +483,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) } } @@ -613,12 +634,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 5a640875..019c7bac 100644 --- a/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/LlamaLanguageModel.swift @@ -706,9 +706,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( @@ -725,12 +730,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 } @@ -852,9 +858,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) @@ -865,18 +868,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, @@ -954,7 +965,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 { @@ -974,11 +989,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) } } } @@ -990,11 +1005,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, @@ -1002,7 +1021,7 @@ import Foundation ) } else { fullPrompt = try formatPrompt( - for: session, + requestContext: requestContext, extraSystemMessage: nil, assistantPrefill: runtimeOptions.assistantPrefill, imageMarker: imageMarker, @@ -1090,18 +1109,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 { @@ -1111,7 +1118,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? @@ -1119,6 +1125,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 @@ -1134,9 +1141,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, @@ -1229,7 +1244,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 @@ -1247,11 +1266,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) @@ -1996,9 +2015,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 { @@ -2131,14 +2150,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, @@ -2148,7 +2167,7 @@ import Foundation } private func formatPrompt( - for session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, extraSystemMessage: String?, assistantPrefill: String?, imageMarker: String?, @@ -2225,7 +2244,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 f9c2eb98..df0da2eb 100644 --- a/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift @@ -1039,8 +1039,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( @@ -1112,10 +1112,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, @@ -1131,8 +1132,6 @@ import Foundation ) } - let toolSpecs = mlxToolSpecs(for: session) - // Map AnyLanguageModel GenerationOptions to MLX GenerateParameters let generateParameters = toGenerateParameters(options) @@ -1142,8 +1141,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] = [] @@ -1154,6 +1152,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, @@ -1210,7 +1215,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 @@ -1233,7 +1238,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 { @@ -1255,7 +1264,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 @@ -1382,11 +1391,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. @@ -1414,6 +1419,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, @@ -1471,7 +1483,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 } @@ -1490,7 +1502,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 { @@ -1506,7 +1522,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() } @@ -1569,11 +1585,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 = try context.model.newCache(parameters: params) @@ -1682,27 +1699,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)) @@ -1894,12 +1911,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 } @@ -2065,7 +2083,7 @@ import Foundation private func generateStructuredJSON( context: ModelContext, tokenCache: StructuredGenerationTokenCache, - session: LanguageModelSession, + requestContext: LanguageModelSession.RequestContext, prompt: Prompt, schema: GenerationSchema, options: GenerationOptions, @@ -2074,7 +2092,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 ae3ce1b2..26e5066d 100644 --- a/Sources/AnyLanguageModel/Models/OllamaLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/OllamaLanguageModel.swift @@ -119,19 +119,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 @@ -140,6 +128,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, @@ -163,7 +157,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 { @@ -252,16 +250,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, @@ -291,14 +292,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, @@ -309,7 +314,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, @@ -347,6 +352,7 @@ private enum ToolResolutionOutcome { private func resolveToolCalls( _ toolCalls: [OllamaToolCall], + tools: [any Tool], session: LanguageModelSession ) async throws -> ToolResolutionOutcome { if toolCalls.isEmpty { @@ -354,7 +360,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 } @@ -680,15 +686,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 ea9a9602..5176e8c9 100644 --- a/Sources/AnyLanguageModel/Models/OpenAILanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/OpenAILanguageModel.swift @@ -468,22 +468,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, @@ -491,8 +478,6 @@ public struct OpenAILanguageModel: LanguageModel { ) case .responses: return try await respondWithResponses( - messages: session.transcript.toOpenAIMessages(), - tools: openAITools, generating: type, schema: schema, options: options, @@ -502,8 +487,6 @@ public struct OpenAILanguageModel: LanguageModel { } private func respondWithChatCompletions( - messages: [OpenAIMessage], - tools: [OpenAITool]?, generating type: Content.Type, schema: GenerationSchema, options: GenerationOptions, @@ -515,11 +498,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, @@ -562,10 +550,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 { @@ -584,7 +576,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)) @@ -621,8 +613,6 @@ public struct OpenAILanguageModel: LanguageModel { } private func respondWithResponses( - messages: [OpenAIMessage], - tools: [OpenAITool]?, generating type: Content.Type, schema: GenerationSchema, options: GenerationOptions, @@ -634,13 +624,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, @@ -669,11 +664,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 { @@ -693,7 +692,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)) @@ -777,16 +776,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 { @@ -835,7 +838,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) } @@ -874,14 +879,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()) @@ -895,7 +906,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)) @@ -1725,12 +1736,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 579bc7cc..b73b4dfd 100644 --- a/Sources/AnyLanguageModel/Models/OpenResponsesLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/OpenResponsesLanguageModel.swift @@ -435,11 +435,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, @@ -489,18 +485,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, @@ -531,7 +530,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) } @@ -548,7 +549,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()) @@ -558,7 +563,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( @@ -583,8 +588,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, @@ -596,11 +599,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, @@ -625,11 +633,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 { @@ -647,7 +661,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)) @@ -1160,11 +1174,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 09804437..d76657c2 100644 --- a/Sources/AnyLanguageModel/Models/PrivateCloudComputeLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/PrivateCloudComputeLanguageModel.swift @@ -126,13 +126,14 @@ issues: [LanguageModelFeedback.Issue], desiredOutput: Transcript.Entry? ) -> Data { + let requestContext = session.resolvedRequestContext() // Attach the feedback to the session's conversation, including its latest response. let fmSession = FoundationModels.LanguageModelSession( model: pccModel, - tools: session.tools.toFoundationModels(), - transcript: session.transcript.toFoundationModels( - instructions: session.instructions, - toolDefinitions: session.tools + tools: requestContext.tools.toFoundationModels(), + transcript: requestContext.transcript.toFoundationModels( + instructions: requestContext.instructions, + toolDefinitions: requestContext.tools .filter(\.includesSchemaInInstructions) .map { Transcript.ToolDefinition(tool: $0) } ) diff --git a/Sources/AnyLanguageModel/Models/SystemLanguageModel.swift b/Sources/AnyLanguageModel/Models/SystemLanguageModel.swift index 0807a2da..8def36c2 100644 --- a/Sources/AnyLanguageModel/Models/SystemLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/SystemLanguageModel.swift @@ -197,19 +197,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, @@ -263,19 +252,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, @@ -290,13 +268,14 @@ issues: [LanguageModelFeedback.Issue], desiredOutput: Transcript.Entry? ) -> Data { + let requestContext = session.resolvedRequestContext() // Attach the feedback to the session's conversation, including its latest response. let fmSession = FoundationModels.LanguageModelSession( model: systemModel, - tools: session.tools.toFoundationModels(), - transcript: session.transcript.toFoundationModels( - instructions: session.instructions, - toolDefinitions: session.tools + tools: requestContext.tools.toFoundationModels(), + transcript: requestContext.transcript.toFoundationModels( + instructions: requestContext.instructions, + toolDefinitions: requestContext.tools .filter(\.includesSchemaInInstructions) .map { Transcript.ToolDefinition(tool: $0) } ) @@ -313,6 +292,26 @@ ) } + private func makeSession( + for session: LanguageModelSession, + prompt: Prompt + ) throws -> FoundationModels.LanguageModelSession { + 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 @@ -332,6 +331,30 @@ return Transcript(entries: transcript.dropLast()) } + #if compiler(>=6.4) && !os(tvOS) + @available(macOS 27.0, iOS 27.0, visionOS 27.0, watchOS 27.0, *) + func makeFoundationModelsSession( + model: Model, + session: LanguageModelSession, + prompt: Prompt + ) -> FoundationModels.LanguageModelSession { + 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/RequestContextTests.swift b/Tests/AnyLanguageModelTests/RequestContextTests.swift new file mode 100644 index 00000000..a39cbf7a --- /dev/null +++ b/Tests/AnyLanguageModelTests/RequestContextTests.swift @@ -0,0 +1,29 @@ +import Testing + +@testable import AnyLanguageModel + +@Suite("Request context") +struct RequestContextTests { + @Test func staticSessionContextMatchesTheSession() { + let session = LanguageModelSession( + model: MockLanguageModel(), + tools: [WeatherTool()], + instructions: "Static" + ) + + let context = session.resolvedRequestContext() + + #expect(context.instructions?.description == "Static") + #expect(context.tools.map(\.name) == session.tools.map(\.name)) + #expect(context.transcript == session.transcript) + } + + @Test func contextReflectsTheCurrentTranscript() async throws { + let session = LanguageModelSession(model: MockLanguageModel.fixed("Hi"), instructions: "Static") + + _ = try await session.respond(to: "Hello") + + #expect(session.resolvedRequestContext().transcript == session.transcript) + #expect(session.resolvedRequestContext().transcript.count == 3) + } +} 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 From d29a850d10ba6988debab0c26f7d721ddf08ee96 Mon Sep 17 00:00:00 2001 From: Mattt Zmuda Date: Sat, 3 Oct 2026 08:21:16 -0700 Subject: [PATCH 2/3] Test that nonstreaming Ollama requests send instructions and history --- .../OllamaRequestHistoryTests.swift | 33 +++++++++++++++++++ 1 file changed, 33 insertions(+) create mode 100644 Tests/AnyLanguageModelTests/OllamaRequestHistoryTests.swift diff --git a/Tests/AnyLanguageModelTests/OllamaRequestHistoryTests.swift b/Tests/AnyLanguageModelTests/OllamaRequestHistoryTests.swift new file mode 100644 index 00000000..51a29632 --- /dev/null +++ b/Tests/AnyLanguageModelTests/OllamaRequestHistoryTests.swift @@ -0,0 +1,33 @@ +import Foundation +import Testing + +@testable import AnyLanguageModel + +#if canImport(Darwin) && !canImport(AsyncHTTPClient) + @Suite("Ollama request history", .serialized) + struct OllamaRequestHistoryTests { + private func response(_ content: String) -> String { + """ + {"model": "test", "created_at": "2026-10-03T00:00:00.000Z", + "message": {"role": "assistant", "content": "\(content)"}, "done": true} + """ + } + + @Test func nonstreamingRequestsSendInstructionsAndHistory() async throws { + ReasoningURLProtocol.reset() + ReasoningURLProtocol.enqueue(json: response("Hi")) + ReasoningURLProtocol.enqueue(json: response("Hi again")) + + let model = OllamaLanguageModel(model: "test", session: ReasoningURLProtocol.makeSession()) + let session = LanguageModelSession(model: model, instructions: "Be brief.") + _ = try await session.respond(to: "Hello") + _ = try await session.respond(to: "Again") + + let body = try #require(ReasoningURLProtocol.recordedBodies.last) + let json = try #require(try JSONSerialization.jsonObject(with: body) as? [String: Any]) + let messages = try #require(json["messages"] as? [[String: Any]]) + #expect(messages.map { $0["role"] as? String } == ["system", "user", "assistant", "user"]) + #expect(messages.map { $0["content"] as? String } == ["Be brief.", "Hello", "Hi", "Again"]) + } + } +#endif From c7a98c05799dea0df885bc5aee3c012cff35a527 Mon Sep 17 00:00:00 2001 From: Mattt Zmuda Date: Sat, 3 Oct 2026 08:28:51 -0700 Subject: [PATCH 3/3] Give the Ollama history test its own URL protocol --- .../OllamaRequestHistoryTests.swift | 74 +++++++++++++++++-- 1 file changed, 69 insertions(+), 5 deletions(-) diff --git a/Tests/AnyLanguageModelTests/OllamaRequestHistoryTests.swift b/Tests/AnyLanguageModelTests/OllamaRequestHistoryTests.swift index 51a29632..4567f203 100644 --- a/Tests/AnyLanguageModelTests/OllamaRequestHistoryTests.swift +++ b/Tests/AnyLanguageModelTests/OllamaRequestHistoryTests.swift @@ -14,20 +14,84 @@ import Testing } @Test func nonstreamingRequestsSendInstructionsAndHistory() async throws { - ReasoningURLProtocol.reset() - ReasoningURLProtocol.enqueue(json: response("Hi")) - ReasoningURLProtocol.enqueue(json: response("Hi again")) + OllamaHistoryURLProtocol.reset(responses: [response("Hi"), response("Hi again")]) - let model = OllamaLanguageModel(model: "test", session: ReasoningURLProtocol.makeSession()) + let model = OllamaLanguageModel(model: "test", session: OllamaHistoryURLProtocol.makeSession()) let session = LanguageModelSession(model: model, instructions: "Be brief.") _ = try await session.respond(to: "Hello") _ = try await session.respond(to: "Again") - let body = try #require(ReasoningURLProtocol.recordedBodies.last) + let body = try #require(OllamaHistoryURLProtocol.recordedBodies.last) let json = try #require(try JSONSerialization.jsonObject(with: body) as? [String: Any]) let messages = try #require(json["messages"] as? [[String: Any]]) #expect(messages.map { $0["role"] as? String } == ["system", "user", "assistant", "user"]) #expect(messages.map { $0["content"] as? String } == ["Be brief.", "Hello", "Hi", "Again"]) } } + + /// A `URLProtocol` with its own queue of JSON responses, + /// so this suite doesn't share state with suites that run in parallel. + private final class OllamaHistoryURLProtocol: URLProtocol { + private struct State: Sendable { + var pending: [String] = [] + var recordedBodies: [Data] = [] + } + + private static let state = Locked(State()) + + static func reset(responses: [String]) { + state.withLock { $0 = State(pending: responses) } + } + + static var recordedBodies: [Data] { + state.withLock { $0.recordedBodies } + } + + static func makeSession() -> URLSession { + let configuration = URLSessionConfiguration.ephemeral + configuration.protocolClasses = [OllamaHistoryURLProtocol.self] + return URLSession(configuration: configuration) + } + + override class func canInit(with request: URLRequest) -> Bool { true } + + override class func canonicalRequest(for request: URLRequest) -> URLRequest { request } + + override func startLoading() { + // URLSession moves `httpBody` to `httpBodyStream` before the protocol sees the request. + let body = request.httpBody ?? request.httpBodyStream.map(Self.readAll) ?? Data() + let next = Self.state.withLock { state -> String? in + state.recordedBodies.append(body) + return state.pending.isEmpty ? nil : state.pending.removeFirst() + } + guard let next, let url = request.url else { + client?.urlProtocol(self, didFailWithError: URLError(.resourceUnavailable)) + return + } + let response = HTTPURLResponse( + url: url, + statusCode: 200, + httpVersion: "HTTP/1.1", + headerFields: ["Content-Type": "application/json"] + )! + client?.urlProtocol(self, didReceive: response, cacheStoragePolicy: .notAllowed) + client?.urlProtocol(self, didLoad: Data(next.utf8)) + client?.urlProtocolDidFinishLoading(self) + } + + override func stopLoading() {} + + private static func readAll(_ stream: InputStream) -> Data { + stream.open() + defer { stream.close() } + var data = Data() + var buffer = [UInt8](repeating: 0, count: 4096) + while true { + let read = stream.read(&buffer, maxLength: buffer.count) + if read <= 0 { break } + data.append(buffer, count: read) + } + return data + } + } #endif