diff --git a/Sources/AnyLanguageModel/GenerationOptions.swift b/Sources/AnyLanguageModel/GenerationOptions.swift index 5d0f6338..f28ab57f 100644 --- a/Sources/AnyLanguageModel/GenerationOptions.swift +++ b/Sources/AnyLanguageModel/GenerationOptions.swift @@ -7,14 +7,14 @@ import JSONSchema /// perform various adjustments on how the model chooses output tokens, /// to specify the penalties for repeating tokens or generating /// longer responses. -public struct GenerationOptions: Sendable, Equatable, Codable { +public struct GenerationOptions: Sendable, Equatable { /// A type that defines how values are sampled from a probability distribution. /// /// A model builds its response to a prompt in a loop. At each iteration in the /// loop the model produces a probability distribution for all the tokens in its /// vocabulary. The sampling mode controls how a token is selected from that /// distribution. - public struct SamplingMode: Sendable, Equatable, Codable { + public struct SamplingMode: Sendable, Equatable { enum Mode: Equatable, Codable { case greedy case topK(Int, seed: UInt64?) @@ -168,6 +168,42 @@ public struct GenerationOptions: Sendable, Equatable, Codable { } } +// MARK: - Transcript Coding + +extension GenerationOptions { + /// The coded form of generation options in a transcript prompt. + /// + /// `GenerationOptions` isn't `Codable`, matching Foundation Models, + /// but ``Transcript/Prompt`` is, so it codes its options through this type. + /// It codes the sampling mode, temperature, and maximum response tokens + /// in the same format as earlier releases. + /// Custom options aren't coded, + /// because decoding them would need a registry of every model's option types. + struct TranscriptCoding: Codable { + struct Sampling: Codable { + var mode: SamplingMode.Mode + } + + var sampling: Sampling? + var temperature: Double? + var maximumResponseTokens: Int? + + init(_ options: GenerationOptions) { + self.sampling = options.sampling.map { Sampling(mode: $0.mode) } + self.temperature = options.temperature + self.maximumResponseTokens = options.maximumResponseTokens + } + + var options: GenerationOptions { + GenerationOptions( + sampling: sampling.map { SamplingMode(mode: $0.mode) }, + temperature: temperature, + maximumResponseTokens: maximumResponseTokens + ) + } + } +} + // MARK: - Custom Generation Options /// A protocol for model-specific generation options. @@ -190,7 +226,7 @@ extension Never: CustomGenerationOptions {} extension Dictionary: CustomGenerationOptions where Key == String, Value == JSONValue {} /// Storage for model-specific custom options. -private struct CustomOptionsStorage: Sendable, Equatable, Codable { +private struct CustomOptionsStorage: Sendable, Equatable { private var storage: [ObjectIdentifier: AnyCustomOptions] = [:] init() {} @@ -222,24 +258,6 @@ private struct CustomOptionsStorage: Sendable, Equatable, Codable { } return true } - - func encode(to encoder: any Encoder) throws { - // Encode custom options that conform to Encodable, keyed by type name - var container = encoder.container(keyedBy: TypeNameCodingKey.self) - for (_, wrapper) in storage { - if let encodeImpl = wrapper.encodeImpl { - let key = TypeNameCodingKey(wrapper.typeName) - let nestedEncoder = container.superEncoder(forKey: key) - try encodeImpl(nestedEncoder) - } - } - } - - init(from decoder: any Decoder) throws { - // Custom options cannot be decoded without a type registry. - // The encoded type names are preserved but the values are lost on round-trip. - self.storage = [:] - } } // MARK: - AnyCustomOptions @@ -247,48 +265,17 @@ private struct CustomOptionsStorage: Sendable, Equatable, Codable { /// A type-erased wrapper for custom generation options. private struct AnyCustomOptions: Sendable { let value: any CustomGenerationOptions - let typeName: String let equalsImpl: @Sendable (any CustomGenerationOptions) -> Bool - let encodeImpl: (@Sendable (any Encoder) throws -> Void)? init(_ value: T) { self.value = value - self.typeName = String(reflecting: T.self) self.equalsImpl = { other in guard let otherTyped = other as? T else { return false } return value == otherTyped } - - // Conditionally capture encode if T conforms to Encodable. - // We capture `value` (which is Sendable) and cast inside the closure. - if value is any Encodable { - self.encodeImpl = { encoder in - // Safe: we checked conformance above, and value is Sendable - try (value as! any Encodable).encode(to: encoder) - } - } else { - self.encodeImpl = nil - } } func isEqual(to other: AnyCustomOptions) -> Bool { equalsImpl(other.value) } } - -private struct TypeNameCodingKey: CodingKey { - var stringValue: String - var intValue: Int? { nil } - - init(_ typeName: String) { - self.stringValue = typeName - } - - init?(stringValue: String) { - self.stringValue = stringValue - } - - init?(intValue: Int) { - nil - } -} diff --git a/Sources/AnyLanguageModel/Transcript.swift b/Sources/AnyLanguageModel/Transcript.swift index 05abfecd..edd6c939 100644 --- a/Sources/AnyLanguageModel/Transcript.swift +++ b/Sources/AnyLanguageModel/Transcript.swift @@ -270,6 +270,31 @@ public struct Transcript: Sendable, Equatable, Codable { self.options = options self.responseFormat = responseFormat } + + private enum CodingKeys: String, CodingKey { + case id, segments, options, responseFormat + } + + public init(from decoder: Decoder) throws { + let container = try decoder.container(keyedBy: CodingKeys.self) + self.id = try container.decode(String.self, forKey: .id) + self.segments = try container.decode([Segment].self, forKey: .segments) + self.options = try container.decode(GenerationOptions.TranscriptCoding.self, forKey: .options).options + self.responseFormat = try container.decodeIfPresent(ResponseFormat.self, forKey: .responseFormat) + } + + /// Encodes this prompt into the given encoder. + /// + /// The encoded ``options`` include the sampling mode, temperature, + /// and maximum response tokens, but not custom options + /// set with ``GenerationOptions/subscript(custom:)``. + public func encode(to encoder: Encoder) throws { + var container = encoder.container(keyedBy: CodingKeys.self) + try container.encode(id, forKey: .id) + try container.encode(segments, forKey: .segments) + try container.encode(GenerationOptions.TranscriptCoding(options), forKey: .options) + try container.encodeIfPresent(responseFormat, forKey: .responseFormat) + } } /// Specifies a response format that the model must conform its output to. diff --git a/Tests/AnyLanguageModelTests/CustomGenerationOptionsTests.swift b/Tests/AnyLanguageModelTests/CustomGenerationOptionsTests.swift index 56f4d7ec..6dddf42e 100644 --- a/Tests/AnyLanguageModelTests/CustomGenerationOptionsTests.swift +++ b/Tests/AnyLanguageModelTests/CustomGenerationOptionsTests.swift @@ -100,47 +100,6 @@ struct CustomGenerationOptionsTests { #expect(options1 != options2) } - - // MARK: - Encoding - - @Test func encodingWithCustomOptions() throws { - var options = GenerationOptions(temperature: 0.8) - options[custom: OpenAILanguageModel.self] = .init( - extraBody: ["reasoning": .object(["enabled": .bool(true)])] - ) - - let encoder = JSONEncoder() - encoder.outputFormatting = [.sortedKeys, .prettyPrinted] - let data = try encoder.encode(options) - let json = String(data: data, encoding: .utf8)! - - // Verify the JSON contains the temperature - #expect(json.contains("\"temperature\"")) - #expect(json.contains("0.8")) - - // Verify custom options type name is in the output - #expect(json.contains("OpenAILanguageModel")) - #expect(json.contains("CustomGenerationOptions")) - } - - @Test func decodingLosesCustomOptions() throws { - var options = GenerationOptions(temperature: 0.8) - options[custom: OpenAILanguageModel.self] = .init( - extraBody: ["key": .string("value")] - ) - - let encoder = JSONEncoder() - let data = try encoder.encode(options) - - let decoder = JSONDecoder() - let decoded = try decoder.decode(GenerationOptions.self, from: data) - - // Standard options should be preserved - #expect(decoded.temperature == 0.8) - - // Custom options are lost on round-trip (documented behavior) - #expect(decoded[custom: OpenAILanguageModel.self] == nil) - } } @Suite("Anthropic CustomGenerationOptions") diff --git a/Tests/AnyLanguageModelTests/TranscriptTests.swift b/Tests/AnyLanguageModelTests/TranscriptTests.swift index 30d5811f..ce7e62d4 100644 --- a/Tests/AnyLanguageModelTests/TranscriptTests.swift +++ b/Tests/AnyLanguageModelTests/TranscriptTests.swift @@ -211,4 +211,50 @@ struct TranscriptTests { } } } + + @Test(arguments: [ + (GenerationOptions.SamplingMode.greedy, #"{"greedy":{}}"#), + (.random(top: 40, seed: 7), #"{"topK":{"_0":40,"seed":7}}"#), + (.random(probabilityThreshold: 0.9), #"{"nucleus":{"_0":0.9}}"#), + ]) + func promptEncodesOptionsWithoutCustomOptions( + sampling: GenerationOptions.SamplingMode, + encodedMode: String + ) throws { + var options = GenerationOptions(sampling: sampling, temperature: 0.5, maximumResponseTokens: 64) + options[custom: OpenAILanguageModel.self] = .init(extraBody: ["key": "value"]) + let prompt = Transcript.Prompt(id: "p", segments: [.text(.init(id: "s", content: "Hi"))], options: options) + + let encoder = JSONEncoder() + encoder.outputFormatting = .sortedKeys + let json = String(decoding: try encoder.encode(prompt), as: UTF8.self) + #expect( + json + == #"{"id":"p","options":{"maximumResponseTokens":64,"sampling":{"mode":"# + + encodedMode + + #"},"temperature":0.5},"segments":[{"text":{"_0":{"content":"Hi","id":"s"}}}]}"# + ) + + let decoded = try JSONDecoder().decode(Transcript.Prompt.self, from: Data(json.utf8)) + #expect(decoded.options == GenerationOptions(sampling: sampling, temperature: 0.5, maximumResponseTokens: 64)) + #expect(decoded.options[custom: OpenAILanguageModel.self] == nil) + } + + @Test func promptDecodesOptionsWithEncodedCustomOptions() throws { + // Earlier releases encoded `GenerationOptions` with a `customOptionsStorage` key. + let json = #""" + {"id":"p","options":{"customOptionsStorage":{"AnyLanguageModel.OpenAILanguageModel.CustomGenerationOptions": + {"extra_body":{"key":"value"}}},"maximumResponseTokens":64,"sampling":{"mode":{"greedy":{}}}, + "temperature":0.5},"segments":[]} + """# + let decoded = try JSONDecoder().decode(Transcript.Prompt.self, from: Data(json.utf8)) + #expect(decoded.options == GenerationOptions(sampling: .greedy, temperature: 0.5, maximumResponseTokens: 64)) + #expect(decoded.options[custom: OpenAILanguageModel.self] == nil) + } + + @Test func promptRoundTripsDefaultOptions() throws { + let prompt = Transcript.Prompt(id: "p", segments: []) + let data = try JSONEncoder().encode(prompt) + #expect(try JSONDecoder().decode(Transcript.Prompt.self, from: data) == prompt) + } }