Skip to content
Open
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
91 changes: 39 additions & 52 deletions Sources/AnyLanguageModel/GenerationOptions.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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?)
Expand Down Expand Up @@ -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.
Expand All @@ -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() {}
Expand Down Expand Up @@ -222,73 +258,24 @@ 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

/// 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<T: CustomGenerationOptions>(_ 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
}
}
25 changes: 25 additions & 0 deletions Sources/AnyLanguageModel/Transcript.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
41 changes: 0 additions & 41 deletions Tests/AnyLanguageModelTests/CustomGenerationOptionsTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
46 changes: 46 additions & 0 deletions Tests/AnyLanguageModelTests/TranscriptTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
Loading