Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
128 changes: 104 additions & 24 deletions Sources/AnyLanguageModel/LanguageModelSession.swift
Original file line number Diff line number Diff line change
@@ -1,13 +1,40 @@
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 {
access(keyPath: \.isResponding)
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 }
Expand All @@ -24,6 +51,18 @@ public final class LanguageModelSession: @unchecked Sendable {

@ObservationIgnored private let state: Locked<State>

@ObservationIgnored private let responseRelays = Locked<[UUID: Task<Void, Never>]>([:])

/// 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?
Expand Down Expand Up @@ -140,13 +179,27 @@ public final class LanguageModelSession: @unchecked Sendable {
}
}

nonisolated private func wrapRespond<T>(_ 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<T>(_ 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
}
Expand All @@ -159,22 +212,33 @@ public final class LanguageModelSession: @unchecked Sendable {
let session = self
let relay = AsyncThrowingStream<ResponseStream<Content>.Snapshot, any Error> { continuation in
let stream = upstream
let task = Task {
session.beginResponding()
var lastSnapshot: ResponseStream<Content>.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<Content>.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 {
Expand All @@ -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 {
Expand Down Expand Up @@ -415,10 +492,11 @@ public final class LanguageModelSession: @unchecked Sendable {
options: GenerationOptions,
generate: () async throws -> Response<Content>
) async throws -> Response<Content> {
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
Expand Down Expand Up @@ -784,7 +862,7 @@ extension LanguageModelSession {
includeSchemaInPrompt: Bool = true,
options: GenerationOptions = GenerationOptions()
) async throws -> Response<Content> where Content: Generable {
try await wrapRespond {
try await wrapRespond { promptID in
// Build segments from text and images
var segments: [Transcript.Segment] = []
if !prompt.isEmpty {
Expand All @@ -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)
Expand Down Expand Up @@ -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 }
Expand Down
Loading
Loading