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
80 changes: 70 additions & 10 deletions Sources/AnyLanguageModel/Generable.swift
Original file line number Diff line number Diff line change
Expand Up @@ -113,10 +113,17 @@ extension Bool: Generable {

/// Creates an instance with the content.
public init(_ content: GeneratedContent) throws {
guard case .bool(let value) = content.kind else {
switch content.kind {
case .bool(let value):
self = value
case .string(let text):
guard let value = Bool(coercing: text) else {
throw GeneratedContentConversionError.typeMismatch
}
self = value
default:
throw GeneratedContentConversionError.typeMismatch
}
self = value
}

/// An instance that represents the generated content.
Expand Down Expand Up @@ -163,10 +170,17 @@ extension Int: Generable {

/// Creates an instance with the content.
public init(_ content: GeneratedContent) throws {
guard case .number(let value) = content.kind else {
switch content.kind {
case .number(let value):
self = Int(value)
case .string(let text):
guard let value = Double(coercing: text), value == value.rounded() else {
throw GeneratedContentConversionError.typeMismatch
}
self = Int(value)
default:
throw GeneratedContentConversionError.typeMismatch
}
self = Int(value)
}

/// An instance that represents the generated content.
Expand All @@ -188,10 +202,17 @@ extension Float: Generable {

/// Creates an instance with the content.
public init(_ content: GeneratedContent) throws {
guard case .number(let value) = content.kind else {
switch content.kind {
case .number(let value):
self = Float(value)
case .string(let text):
guard let value = Double(coercing: text) else {
throw GeneratedContentConversionError.typeMismatch
}
self = Float(value)
default:
throw GeneratedContentConversionError.typeMismatch
}
self = Float(value)
}

/// An instance that represents the generated content.
Expand All @@ -213,10 +234,17 @@ extension Double: Generable {

/// Creates an instance with the content.
public init(_ content: GeneratedContent) throws {
guard case .number(let value) = content.kind else {
switch content.kind {
case .number(let value):
self = value
case .string(let text):
guard let value = Double(coercing: text) else {
throw GeneratedContentConversionError.typeMismatch
}
self = value
default:
throw GeneratedContentConversionError.typeMismatch
}
self = value
}

/// An instance that represents the generated content.
Expand All @@ -238,10 +266,17 @@ extension Decimal: Generable {

/// Creates an instance with the content.
public init(_ content: GeneratedContent) throws {
guard case .number(let value) = content.kind else {
switch content.kind {
case .number(let value):
self = Decimal(value)
case .string(let text):
guard let value = Double(coercing: text) else {
throw GeneratedContentConversionError.typeMismatch
}
self = Decimal(value)
default:
throw GeneratedContentConversionError.typeMismatch
}
self = Decimal(value)
}

/// An instance that represents the generated content.
Expand All @@ -251,6 +286,31 @@ extension Decimal: Generable {
}
}

// MARK: String Coercion

/// Language models sometimes emit primitive values as JSON strings, most commonly in
/// tool-call arguments where a chat template stringifies every argument value. The
/// primitive initializers above accept such strings when they parse unambiguously as
/// the target type, and throw ``GeneratedContentConversionError/typeMismatch`` otherwise.
extension Bool {
fileprivate init?(coercing text: String) {
switch text.trimmingCharacters(in: .whitespacesAndNewlines).lowercased() {
case "true":
self = true
case "false":
self = false
default:
return nil
}
}
}

extension Double {
fileprivate init?(coercing text: String) {
self.init(text.trimmingCharacters(in: .whitespacesAndNewlines))
}
}

// MARK: Array

extension Array: Generable where Element: Generable {
Expand Down
56 changes: 56 additions & 0 deletions Tests/AnyLanguageModelTests/PrimitiveStringCoercionTests.swift
Original file line number Diff line number Diff line change
@@ -0,0 +1,56 @@
import Foundation
import Testing

@testable import AnyLanguageModel

@Generable
private struct StringifiedToolArguments {
var count: Int?
var enabled: Bool?
var ratio: Double?
}

@Suite("Primitive string coercion")
struct PrimitiveStringCoercionTests {
@Test func decodesStringifiedNumbersAndBools() throws {
#expect(try Int(GeneratedContent(kind: .string("5"))) == 5)
#expect(try Int(GeneratedContent(kind: .string(" -12 "))) == -12)
#expect(try Int(GeneratedContent(kind: .string("5.0"))) == 5)
#expect(try Bool(GeneratedContent(kind: .string("true"))) == true)
#expect(try Bool(GeneratedContent(kind: .string("False"))) == false)
#expect(try Double(GeneratedContent(kind: .string("3.25"))) == 3.25)
#expect(try Float(GeneratedContent(kind: .string("0.5"))) == 0.5)
#expect(try Decimal(GeneratedContent(kind: .string("2.5"))) == Decimal(2.5))
}

@Test func rejectsUnparseableStrings() {
#expect(throws: GeneratedContentConversionError.self) {
try Int(GeneratedContent(kind: .string("five")))
}
#expect(throws: GeneratedContentConversionError.self) {
try Int(GeneratedContent(kind: .string("5.5")))
}
#expect(throws: GeneratedContentConversionError.self) {
try Bool(GeneratedContent(kind: .string("yes")))
}
#expect(throws: GeneratedContentConversionError.self) {
try Double(GeneratedContent(kind: .string("")))
}
}

@Test func nativeKindsStillDecode() throws {
#expect(try Int(GeneratedContent(kind: .number(7))) == 7)
#expect(try Bool(GeneratedContent(kind: .bool(true))) == true)
#expect(try Double(GeneratedContent(kind: .number(1.5))) == 1.5)
#expect(try String(GeneratedContent(kind: .string("text"))) == "text")
}

@Test func decodesGenerableArgumentsWithStringifiedValues() throws {
let json = #"{"count": "3", "enabled": "true", "ratio": "0.75"}"#
let arguments = try StringifiedToolArguments(GeneratedContent(json: json))

#expect(arguments.count == 3)
#expect(arguments.enabled == true)
#expect(arguments.ratio == 0.75)
}
}