diff --git a/Sources/AnyLanguageModel/LanguageModelSession.swift b/Sources/AnyLanguageModel/LanguageModelSession.swift index 0a15533a..f027b894 100644 --- a/Sources/AnyLanguageModel/LanguageModelSession.swift +++ b/Sources/AnyLanguageModel/LanguageModelSession.swift @@ -1,6 +1,18 @@ import Foundation import Observation +/// Controls transcript retention when generation fails or is cancelled. +public struct TranscriptErrorHandlingPolicy: Sendable, Equatable { + private let shouldRevert: Bool + + /// Retain the prompt and the latest cumulative streaming checkpoint. + /// Nonstreaming failures have no checkpoint and retain only the prompt. + public static let preserveTranscript = Self(shouldRevert: false) + /// Remove entries added by the failing request, preserving other requests. + /// Tool side effects are not undone. + public static let revertTranscript = Self(shouldRevert: true) +} + @Observable public final class LanguageModelSession: @unchecked Sendable { public var isResponding: Bool { @@ -8,6 +20,21 @@ public final class LanguageModelSession: @unchecked Sendable { return state.withLock { $0.isResponding } } + /// On failure, `nil` retains the prompt without committing streaming checkpoints. + /// Streaming cancellation never commits the partial answer as a completed response, regardless + /// of this policy. These defaults do not claim Foundation Models behavioral parity. + public var transcriptErrorHandlingPolicy: TranscriptErrorHandlingPolicy? { + get { + access(keyPath: \.transcriptErrorHandlingPolicy) + return state.withLock { $0.transcriptErrorHandlingPolicy } + } + set { + withMutation(keyPath: \.transcriptErrorHandlingPolicy) { + state.withLock { $0.transcriptErrorHandlingPolicy = newValue } + } + } + } + public var transcript: Transcript { access(keyPath: \.transcript) return state.withLock { $0.transcript } @@ -24,6 +51,18 @@ public final class LanguageModelSession: @unchecked Sendable { @ObservationIgnored private let state: Locked + @ObservationIgnored private let responseRelays = Locked<[UUID: Task]>([:]) + + /// Waits for transcript cleanup of all streaming relays registered when this call begins. + /// Relays started later are excluded. Cancelling this wait does not cancel generation. + /// Call after cancelling consumers and before persisting the transcript. This is an + /// AnyLanguageModel extension; nonstreaming operations must be awaited separately. + /// Do not call from a tool executing within one of the included responses. + nonisolated public func waitForResponseCompletion() async { + let tasks = responseRelays.withLock { Array($0.values) } + for task in tasks { await task.value } + } + private let model: any LanguageModel public let tools: [any Tool] public let instructions: Instructions? @@ -140,13 +179,27 @@ public final class LanguageModelSession: @unchecked Sendable { } } - nonisolated private func wrapRespond(_ operation: () async throws -> T) async throws -> T { + // The prompt is the only entry committed before a request succeeds. Checkpoints and + // the final response commit atomically at completion, so rollback owns only this ID. + nonisolated private func removePrompt(id: String) { + withMutation(keyPath: \.transcript) { + state.withLock { state in + state.transcript = Transcript(entries: state.transcript.filter { $0.id != id }) + } + } + } + + nonisolated private func wrapRespond(_ operation: (String) async throws -> T) async throws -> T { + let promptID = UUID().uuidString beginResponding() do { - let result = try await operation() + let result = try await operation(promptID) endResponding() return result } catch { + if transcriptErrorHandlingPolicy == .revertTranscript { + removePrompt(id: promptID) + } endResponding() throw error } @@ -159,22 +212,33 @@ public final class LanguageModelSession: @unchecked Sendable { let session = self let relay = AsyncThrowingStream.Snapshot, any Error> { continuation in let stream = upstream - let task = Task { - session.beginResponding() - var lastSnapshot: ResponseStream.Snapshot? - var accountedUsage = Usage.zero - do { - for try await snapshot in stream { - lastSnapshot = snapshot - session.recordUsage(snapshot.usage.increment(since: &accountedUsage)) - continuation.yield(snapshot) - } + let relayID = UUID() + // Publish the task under the same lock used by completion removal, so a + // fast relay cannot finish before registration and leave a stale handle. + let task = session.responseRelays.withLock { relays in + let task = Task { + defer { session.responseRelays.withLock { _ = $0.removeValue(forKey: relayID) } } + session.beginResponding() + var lastSnapshot: ResponseStream.Snapshot? + var accountedUsage = Usage.zero + do { + for try await snapshot in stream { + lastSnapshot = snapshot + session.recordUsage(snapshot.usage.increment(since: &accountedUsage)) + continuation.yield(snapshot) + } - // Commit the response to the transcript - // before the stream reports completion, - // so a caller that drains the stream - // and starts the next turn sees the full history. - if let lastSnapshot { + // AsyncThrowingStream may end normally when its task is cancelled. + // A partial snapshot must not become a completed transcript response. + try Task.checkCancellation() + + // Commit the response to the transcript + // before the stream reports completion, + // so a caller that drains the stream + // and starts the next turn sees the full history. + guard let lastSnapshot else { + throw ResponseStreamError.noSnapshots + } // Extract text content from the generated content let textContent: String if case .string(let str) = lastSnapshot.rawContent.kind { @@ -196,13 +260,26 @@ public final class LanguageModelSession: @unchecked Sendable { $0.transcript.append(responseEntry) } } + session.endResponding() + continuation.finish() + } catch { + session.withMutation(keyPath: \.transcript) { + session.state.withLock { state in + if state.transcriptErrorHandlingPolicy == .preserveTranscript { + state.transcript.append(contentsOf: lastSnapshot?.transcriptEntries ?? []) + } else if state.transcriptErrorHandlingPolicy == .revertTranscript { + state.transcript = Transcript( + entries: state.transcript.filter { $0.id != promptEntry.id } + ) + } + } + } + session.endResponding() + continuation.finish(throwing: error) } - session.endResponding() - continuation.finish() - } catch { - session.endResponding() - continuation.finish(throwing: error) } + relays[relayID] = task + return task } continuation.onTermination = { termination in if case .cancelled = termination { @@ -415,10 +492,11 @@ public final class LanguageModelSession: @unchecked Sendable { options: GenerationOptions, generate: () async throws -> Response ) async throws -> Response { - try await wrapRespond { + try await wrapRespond { promptID in // Add prompt to transcript let promptEntry = Transcript.Entry.prompt( Transcript.Prompt( + id: promptID, segments: [.text(.init(content: prompt.description))], options: options, responseFormat: responseFormat @@ -784,7 +862,7 @@ extension LanguageModelSession { includeSchemaInPrompt: Bool = true, options: GenerationOptions = GenerationOptions() ) async throws -> Response where Content: Generable { - try await wrapRespond { + try await wrapRespond { promptID in // Build segments from text and images var segments: [Transcript.Segment] = [] if !prompt.isEmpty { @@ -795,6 +873,7 @@ extension LanguageModelSession { // Add prompt to transcript let promptEntry = Transcript.Entry.prompt( Transcript.Prompt( + id: promptID, segments: segments, options: options, responseFormat: type == String.self ? nil : .init(type: type) @@ -1235,6 +1314,7 @@ private enum ResponseStreamError: Error, LocalizedError { private struct State: Equatable, Sendable { var transcript: Transcript + var transcriptErrorHandlingPolicy: TranscriptErrorHandlingPolicy? var usage = LanguageModelSession.Usage.zero var isResponding: Bool { count > 0 } diff --git a/Tests/AnyLanguageModelTests/TranscriptErrorHandlingTests.swift b/Tests/AnyLanguageModelTests/TranscriptErrorHandlingTests.swift new file mode 100644 index 00000000..e6ea4dc3 --- /dev/null +++ b/Tests/AnyLanguageModelTests/TranscriptErrorHandlingTests.swift @@ -0,0 +1,237 @@ +import Foundation +import Testing + +@testable import AnyLanguageModel + +@Suite("Transcript error handling") +struct TranscriptErrorHandlingTests { + @Test(arguments: [nil, .preserveTranscript, .revertTranscript] as [TranscriptErrorHandlingPolicy?], [false, true]) + func emptyStreamAppliesPolicyBeforeReportingFailure(policy: TranscriptErrorHandlingPolicy?, collect: Bool) + async throws + { + let previous = Transcript.Entry.prompt(.init(id: "previous", segments: [.text(.init(content: "Previous"))])) + let session = LanguageModelSession( + model: CheckpointModel(completedTools: false, empty: true), + transcript: Transcript(entries: [previous]) + ) + session.transcriptErrorHandlingPolicy = policy + do { + let stream = session.streamResponse(to: "Empty request") + if collect { + _ = try await stream.collect() + } else { + for try await _ in stream { Issue.record("Empty stream yielded a snapshot") } + } + Issue.record("Expected noSnapshots error") + } catch { + #expect(String(describing: error) == "noSnapshots") + // Cleanup must have run before the error reaches either kind of consumer. + #expect(!session.isResponding) + if policy == .revertTranscript { + #expect(Array(session.transcript) == [previous]) + } else { + #expect(session.transcript.count == 2) + #expect(session.transcript.first == previous) + guard case .prompt(let prompt) = session.transcript.last else { + Issue.record("Expected retained request prompt"); return + } + #expect(prompt.segments.first?.description == "Empty request") + } + #expect(!session.transcript.contains { if case .response = $0 { true } else { false } }) + } + await session.waitForResponseCompletion() + } + + @Test(arguments: [false, true]) + func cancelledPartialAnswerIsNotCommitted(completedTools: Bool) async throws { + let session = LanguageModelSession(model: CheckpointModel(completedTools: completedTools)) + session.transcriptErrorHandlingPolicy = .preserveTranscript + let consumer = Task { + for try await _ in session.streamResponse(to: "Question") { + withUnsafeCurrentTask { $0?.cancel() } + } + } + _ = await consumer.result + await session.waitForResponseCompletion() + #expect(!session.isResponding) + #expect(session.transcript.count == (completedTools ? 3 : 1)) + #expect(!session.transcript.contains { if case .response = $0 { true } else { false } }) + if completedTools { + #expect(session.transcript[1].id == "calls") + #expect(session.transcript[2].id == "call") + } + } + + @Test func nilPolicyRetainsPromptWithoutCommittingPartialAnswerOrTools() async throws { + let session = LanguageModelSession(model: CheckpointModel(completedTools: true)) + #expect(session.transcriptErrorHandlingPolicy == nil) + let consumer = Task { + for try await _ in session.streamResponse(to: "Question") { + withUnsafeCurrentTask { $0?.cancel() } + } + } + _ = await consumer.result + await session.waitForResponseCompletion() + #expect(session.transcript.count == 1) + guard case .prompt = session.transcript[0] else { Issue.record("Expected prompt only"); return } + } + + @Test func revertCancellationRestoresPreviousTranscript() async throws { + let old = Transcript.Entry.prompt(.init(id: "previous", segments: [.text(.init(content: "Previous"))])) + let session = LanguageModelSession( + model: CheckpointModel(completedTools: true), + transcript: Transcript(entries: [old]) + ) + session.transcriptErrorHandlingPolicy = .revertTranscript + let consumer = Task { + for try await _ in session.streamResponse(to: "Question") { + withUnsafeCurrentTask { $0?.cancel() } + } + } + _ = await consumer.result + await session.waitForResponseCompletion() + #expect(Array(session.transcript) == [old]) + } + + @Test(arguments: [TranscriptErrorHandlingPolicy.preserveTranscript, .revertTranscript]) + func thrownFailureUsesLatestCumulativeCheckpoint(policy: TranscriptErrorHandlingPolicy) async throws { + let session = LanguageModelSession(model: CheckpointModel(completedTools: true, finish: .failure)) + session.transcriptErrorHandlingPolicy = policy + await #expect(throws: FixtureError.self) { + for try await _ in session.streamResponse(to: "Question") {} + } + await session.waitForResponseCompletion() + if policy == .preserveTranscript { + #expect(session.transcript.count == 3) + #expect(session.transcript.filter { if case .toolOutput = $0 { true } else { false } }.count == 1) + } else { + #expect(session.transcript.isEmpty) + } + #expect(!session.isResponding) + #expect(!session.transcript.contains { if case .response = $0 { true } else { false } }) + } + + @Test func nilPolicyThrownFailureRetainsPrompt() async throws { + let session = LanguageModelSession(model: CheckpointModel(completedTools: true, finish: .failure)) + await #expect(throws: FixtureError.self) { + for try await _ in session.streamResponse(to: "Question") {} + } + #expect(session.transcript.count == 1) + } + + @Test func nonstreamFailureRevertsPrompt() async throws { + let previous = Transcript.Entry.prompt(.init(id: "previous", segments: [.text(.init(content: "Previous"))])) + let session = LanguageModelSession( + model: CheckpointModel(completedTools: true, finish: .failure), + transcript: Transcript(entries: [previous]) + ) + session.transcriptErrorHandlingPolicy = .revertTranscript + await #expect(throws: FixtureError.self) { _ = try await session.respond(to: "Question") } + #expect(Array(session.transcript) == [previous]) + #expect(!session.isResponding) + } + + @Test func nonstreamPreserveFailureHasNoCheckpointChannel() async throws { + let session = LanguageModelSession(model: CheckpointModel(completedTools: true, finish: .failure)) + session.transcriptErrorHandlingPolicy = .preserveTranscript + await #expect(throws: FixtureError.self) { _ = try await session.respond(to: "Question") } + #expect(session.transcript.count == 1) + } + + @Test(arguments: [TranscriptErrorHandlingPolicy.preserveTranscript, .revertTranscript]) + func successStillCommitsFullResponseOnce(policy: TranscriptErrorHandlingPolicy) async throws { + let session = LanguageModelSession(model: CheckpointModel(completedTools: true, finish: .success)) + session.transcriptErrorHandlingPolicy = policy + let result = try await session.streamResponse(to: "Question").collect() + await session.waitForResponseCompletion() + #expect(result.content == "Answer") + #expect(session.transcript.count == 4) + #expect(session.transcript.filter { if case .response = $0 { true } else { false } }.count == 1) + #expect(session.transcript.filter { if case .toolOutput = $0 { true } else { false } }.count == 1) + } +} + +private enum FixtureError: Error { case failed } + +private struct CheckpointModel: LanguageModel { + typealias UnavailableReason = Never + enum Finish: Sendable { case suspended, failure, success } + var completedTools: Bool + var finish: Finish = .suspended + var empty = false + + func respond( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) async throws -> LanguageModelSession.Response { + try await streamResponse( + within: session, + to: prompt, + generating: type, + includeSchemaInPrompt: includeSchemaInPrompt, + options: options + ).collect() + } + + func streamResponse( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) -> sending LanguageModelSession.ResponseStream { + .init( + stream: AsyncThrowingStream { continuation in + if empty { continuation.finish(); return } + do { + var entries: [Transcript.Entry] = [] + if completedTools { + entries.append( + .toolCalls( + .init( + id: "calls", + [ + .init(id: "call", toolName: "fixture", arguments: GeneratedContent("{}")) + ] + ) + ) + ) + entries.append( + .toolOutput( + .init(id: "call", toolName: "fixture", segments: [.text(.init(content: "Done"))]) + ) + ) + } + // Repeated cumulative entries exercise replacement/checkpoint behavior. + for text in ["Partial answer", "Partial answer"] { + let raw = GeneratedContent(text) + continuation.yield( + .init( + content: try Content(raw).asPartiallyGenerated(), + rawContent: raw, + transcriptEntries: ArraySlice(entries) + ) + ) + } + switch finish { + case .suspended: return + case .failure: throw FixtureError.failed + case .success: + let raw = GeneratedContent("Answer") + continuation.yield( + .init( + content: try Content(raw).asPartiallyGenerated(), + rawContent: raw, + transcriptEntries: ArraySlice(entries) + ) + ) + continuation.finish() + } + } catch { continuation.finish(throwing: error) } + } + ) + } +} diff --git a/Tests/AnyLanguageModelTests/TranscriptOverlapTests.swift b/Tests/AnyLanguageModelTests/TranscriptOverlapTests.swift new file mode 100644 index 00000000..1fefb0c1 --- /dev/null +++ b/Tests/AnyLanguageModelTests/TranscriptOverlapTests.swift @@ -0,0 +1,229 @@ +import Foundation +import Observation +import Testing + +@testable import AnyLanguageModel + +@Suite("Overlapping transcript requests") +struct TranscriptOverlapTests { + @Test(arguments: [ + (false, false, false), (false, false, true), (false, true, false), (false, true, true), + (true, false, false), (true, false, true), (true, true, false), (true, true, true), + ]) + func rollbackKeepsOtherRequest(scenario: (Bool, Bool, Bool)) async throws { + let (streaming, failingStartsFirst, failingFinishesFirst) = scenario + let failing = ControlledRequest(fails: true) + let successful = ControlledRequest(fails: false) + let session = LanguageModelSession(model: OverlapModel(failing: failing, successful: successful)) + session.transcriptErrorHandlingPolicy = .revertTranscript + func start(_ prompt: String) -> Task { + Task { + if streaming { + _ = try await session.streamResponse(to: prompt).collect() + } else { + _ = try await session.respond(to: prompt) + } + } + } + let first = start(failingStartsFirst ? "Fail" : "Keep") + await (failingStartsFirst ? failing : successful).entered.wait() + let second = start(failingStartsFirst ? "Keep" : "Fail") + await (failingStartsFirst ? successful : failing).entered.wait() + let failedTask = failingStartsFirst ? first : second + let keptTask = failingStartsFirst ? second : first + if failingFinishesFirst { + await failing.release.open() + _ = await failedTask.result + await successful.release.open() + try await keptTask.value + } else { + await successful.release.open() + try await keptTask.value + await failing.release.open() + _ = await failedTask.result + } + await session.waitForResponseCompletion() + #expect(!session.isResponding) + #expect(session.transcript.count == 4) + #expect( + session.transcript.contains { + if case .prompt(let p) = $0 { p.segments.first?.description == "Keep" } else { false } + } + ) + #expect( + session.transcript.contains { + if case .response(let r) = $0 { r.segments.first?.description == "Kept answer" } else { false } + } + ) + #expect(session.transcript.contains { $0.id == "kept-calls" }) + #expect(session.transcript.contains { $0.id == "kept-output" }) + } + + @Test func multimodalRollbackKeepsOverlappingResponse() async throws { + let failing = ControlledRequest(fails: true) + let successful = ControlledRequest(fails: false) + let session = LanguageModelSession(model: OverlapModel(failing: failing, successful: successful)) + session.transcriptErrorHandlingPolicy = .revertTranscript + let failure = Task { try await session.respond(to: "Fail", images: [], generating: String.self) } + await failing.entered.wait() + let kept = Task { try await session.respond(to: "Keep") } + await successful.entered.wait() + await successful.release.open() + _ = try await kept.value + await failing.release.open() + _ = await failure.result + #expect(session.transcript.count == 4) + #expect(session.transcript.contains { $0.id == "kept-output" }) + } + + @Test func cancelledOlderRelayKeepsCompletedNewerRequest() async throws { + let older = ControlledRequest(fails: true) + let newer = ControlledRequest(fails: false) + let session = LanguageModelSession(model: OverlapModel(failing: older, successful: newer)) + session.transcriptErrorHandlingPolicy = .revertTranscript + let first = Task { try await session.streamResponse(to: "Fail").collect() } + await older.entered.wait() + let second = Task { try await session.streamResponse(to: "Keep").collect() } + await newer.entered.wait() + await newer.release.open() + _ = try await second.value + first.cancel() + _ = await first.result + await session.waitForResponseCompletion() + #expect(!session.isResponding) + #expect(session.transcript.count == 4) + #expect(session.transcript.contains { $0.id == "kept-output" }) + await older.release.open() + } + + @Test func policyChangesNotifyObservation() { + let session = LanguageModelSession( + model: OverlapModel(failing: .init(fails: true), successful: .init(fails: false)) + ) + let changes = Locked(0) + withObservationTracking { + #expect(session.transcriptErrorHandlingPolicy == nil) + } onChange: { + changes.withLock { $0 += 1 } + } + session.transcriptErrorHandlingPolicy = .preserveTranscript + #expect(changes.withLock { $0 } == 1) + withObservationTracking { + #expect(session.transcriptErrorHandlingPolicy == .preserveTranscript) + } onChange: { + changes.withLock { $0 += 1 } + } + session.transcriptErrorHandlingPolicy = .revertTranscript + #expect(changes.withLock { $0 } == 2) + } + + @Test func completionWaitIncludesOlderStillRunningRelay() async throws { + let older = ControlledRequest(fails: true) + let newer = ControlledRequest(fails: false) + let session = LanguageModelSession(model: OverlapModel(failing: older, successful: newer)) + session.transcriptErrorHandlingPolicy = .revertTranscript + let first = Task { try await session.streamResponse(to: "Fail").collect() } + await older.entered.wait() + let second = Task { try await session.streamResponse(to: "Keep").collect() } + await newer.entered.wait() + await newer.release.open() + _ = try await second.value + // The newer relay is already complete, but the older one still owns an active prompt. + #expect(session.isResponding) + let waiter = Task { + await session.waitForResponseCompletion() + #expect(!session.isResponding) + #expect(session.transcript.count == 4) + } + await older.release.open() + _ = await first.result + await waiter.value + } +} + +private actor RequestGate { + private var isOpen = false + private var waiters: [CheckedContinuation] = [] + func wait() async { + if isOpen { return } + await withCheckedContinuation { waiters.append($0) } + } + func open() { + isOpen = true + let pending = waiters + waiters.removeAll() + for waiter in pending { waiter.resume() } + } +} + +private struct ControlledRequest: Sendable { + let fails: Bool + let entered = RequestGate() + let release = RequestGate() +} + +private enum OverlapError: Error { case expected } + +private struct OverlapModel: LanguageModel { + typealias UnavailableReason = Never + let failing: ControlledRequest + let successful: ControlledRequest + + func respond( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) async throws -> LanguageModelSession.Response { + let request = prompt.description == "Fail" ? failing : successful + await request.entered.open() + await request.release.wait() + if request.fails { throw OverlapError.expected } + let raw = GeneratedContent("Kept answer") + return .init(content: try Content(raw), rawContent: raw, transcriptEntries: Self.entries) + } + + func streamResponse( + within session: LanguageModelSession, + to prompt: Prompt, + generating type: Content.Type, + includeSchemaInPrompt: Bool, + options: GenerationOptions + ) -> sending LanguageModelSession.ResponseStream { + let request = prompt.description == "Fail" ? failing : successful + return .init( + stream: AsyncThrowingStream { continuation in + let task = Task { + await request.entered.open() + await request.release.wait() + do { + if request.fails { throw OverlapError.expected } + let raw = GeneratedContent("Kept answer") + continuation.yield( + .init( + content: try Content(raw).asPartiallyGenerated(), + rawContent: raw, + transcriptEntries: Self.entries + ) + ) + continuation.finish() + } catch { continuation.finish(throwing: error) } + } + continuation.onTermination = { _ in task.cancel() } + } + ) + } + + private static var entries: ArraySlice { + [ + .toolCalls( + .init( + id: "kept-calls", + [.init(id: "kept-output", toolName: "fixture", arguments: GeneratedContent("{}"))] + ) + ), + .toolOutput(.init(id: "kept-output", toolName: "fixture", segments: [.text(.init(content: "Done"))])), + ] + } +}