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
4 changes: 2 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -576,8 +576,8 @@ say which API they follow.
[observing and controlling tool calls](#tool-calling).
- `transcriptErrorHandlingPolicy` and `waitForResponseCompletion()`:
what a transcript keeps when a request fails or is cancelled.
- `LanguageModelSession.tools` and `instructions`:
the session's tools and instructions,
- `LanguageModelSession.tools`, `instructions`, and `resolvedRequestContext()`:
the session's tools, instructions, and the inputs for each request,
for language models defined outside AnyLanguageModel.
- `Usage` and the `usage` properties:
[token usage](#token-usage),
Expand Down
43 changes: 43 additions & 0 deletions Sources/AnyLanguageModel/LanguageModelSession.swift
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,49 @@ public final class LanguageModelSession: @unchecked Sendable {
/// with the Foundation Models framework.
@ObservationIgnored public var toolExecutionDelegate: (any ToolExecutionDelegate)?

/// The transcript, instructions, and tools for one model request.
///
/// A language model creates one context immediately before each request it sends,
/// by calling ``LanguageModelSession/resolvedRequestContext()``.
/// If that request produces tool calls,
/// run them with the ``tools`` from the same context,
/// and resolve a new context only before the continuation request.
///
/// - Note: This API is exclusive to AnyLanguageModel
/// and using it means your code is no longer drop-in compatible
/// with the Foundation Models framework.
/// It's public so that language models outside this module
/// can read the inputs for each request.
public struct RequestContext: Sendable {
/// The transcript to send with the request.
public let transcript: Transcript

/// The instructions for the request, if any.
public let instructions: Instructions?

/// The tools that the model can call in response to the request.
public let tools: [any Tool]

fileprivate init(transcript: Transcript, instructions: Instructions?, tools: [any Tool]) {
self.transcript = transcript
self.instructions = instructions
self.tools = tools
}
}

/// Returns the transcript, instructions, and tools for the next model request.
///
/// Calling this method doesn't change the session.
///
/// - Note: This API is exclusive to AnyLanguageModel
/// and using it means your code is no longer drop-in compatible
/// with the Foundation Models framework.
/// It's public so that language models outside this module
/// can read the inputs for each request.
nonisolated public func resolvedRequestContext() -> RequestContext {
RequestContext(transcript: transcript, instructions: instructions, tools: tools)
}

/// Creates a session with a model, tools,
/// and instructions that you build with a result builder.
///
Expand Down
37 changes: 23 additions & 14 deletions Sources/AnyLanguageModel/Models/AnthropicLanguageModel.swift
Original file line number Diff line number Diff line change
Expand Up @@ -445,17 +445,18 @@ public struct AnthropicLanguageModel: LanguageModel {
) async throws -> LanguageModelSession.Response<Content> where Content: Generable {
let url = baseURL.appendingPathComponent("v1/messages")
let headers = buildHeaders()
let requestContext = session.resolvedRequestContext()

// Convert available tools to Anthropic format
let anthropicTools: [AnthropicTool] = try session.tools.map { tool in
let anthropicTools: [AnthropicTool] = try requestContext.tools.map { tool in
try convertToolToAnthropicFormat(tool)
}

let responseSchema = type == String.self ? nil : try convertSchemaToAnthropicFormat(schema)
let params = try createMessageParams(
model: model,
system: nil,
messages: try session.transcript.toAnthropicMessages(),
messages: try requestContext.transcript.toAnthropicMessages(),
tools: anthropicTools.isEmpty ? nil : anthropicTools,
responseSchema: responseSchema,
options: options
Expand Down Expand Up @@ -486,7 +487,11 @@ public struct AnthropicLanguageModel: LanguageModel {
}

if !toolUses.isEmpty {
let resolution = try await resolveToolUses(toolUses, session: session)
let resolution = try await resolveToolUses(
toolUses,
tools: requestContext.tools,
session: session
)
switch resolution {
case .stop(let calls):
if !calls.isEmpty {
Expand Down Expand Up @@ -584,19 +589,18 @@ public struct AnthropicLanguageModel: LanguageModel {
let task = Task { @Sendable in
do {
let headers = buildHeaders()

// Convert available tools to Anthropic format
let anthropicTools: [AnthropicTool] = try session.tools.map { tool in
try convertToolToAnthropicFormat(tool)
}

let responseSchema =
type == String.self ? nil : try convertSchemaToAnthropicFormat(schema)
var messages = try session.transcript.toAnthropicMessages()
var inFlightMessages: [AnthropicMessage] = []
var state = StreamingResponseState<Content>()
var toolRounds = ToolRoundLimit(provider: "Anthropic")
while true {
try Task.checkCancellation()
let requestContext = session.resolvedRequestContext()
let anthropicTools: [AnthropicTool] = try requestContext.tools.map {
try convertToolToAnthropicFormat($0)
}
let messages = try requestContext.transcript.toAnthropicMessages() + inFlightMessages
var params = try createMessageParams(
model: model,
system: nil,
Expand Down Expand Up @@ -692,14 +696,18 @@ public struct AnthropicLanguageModel: LanguageModel {
guard !toolUses.isEmpty else { break }
try Task.checkCancellation()
try toolRounds.record(toolUses.map(\.roundCall))
switch try await resolveToolUses(toolUses, session: session) {
switch try await resolveToolUses(
toolUses,
tools: requestContext.tools,
session: session
) {
case .stop(let calls):
state.entries.append(.toolCalls(Transcript.ToolCalls(calls)))
continuation.yield(try state.stoppedSnapshot())
continuation.finish()
return
case .invocations(let invocations):
messages.append(.init(role: .assistant, content: content))
inFlightMessages.append(.init(role: .assistant, content: content))
state.entries.append(.toolCalls(Transcript.ToolCalls(invocations.map(\.call))))
var results: [AnthropicContent] = []
for invocation in invocations {
Expand All @@ -713,7 +721,7 @@ public struct AnthropicLanguageModel: LanguageModel {
)
)
}
messages.append(.init(role: .user, content: results))
inFlightMessages.append(.init(role: .user, content: results))
}
if let snapshot = snapshot() { continuation.yield(snapshot) }
state.beginNextRound()
Expand Down Expand Up @@ -893,12 +901,13 @@ private func convertSchemaToAnthropicFormat(_ schema: GenerationSchema) throws -

private func resolveToolUses(
_ toolUses: [AnthropicToolUse],
tools: [any Tool],
session: LanguageModelSession
) async throws -> ToolResolutionOutcome {
if toolUses.isEmpty { return .invocations([]) }

var toolsByName: [String: any Tool] = [:]
for tool in session.tools {
for tool in tools {
if toolsByName[tool.name] == nil {
toolsByName[tool.name] = tool
}
Expand Down
40 changes: 16 additions & 24 deletions Sources/AnyLanguageModel/Models/CoreMLLanguageModel.swift
Original file line number Diff line number Diff line change
Expand Up @@ -113,11 +113,12 @@
includeSchemaInPrompt: Bool,
options: GenerationOptions
) async throws -> LanguageModelSession.Response<Content> where Content: Generable {
try validateNoImageSegments(in: session)
let requestContext = session.resolvedRequestContext()
try validateNoImageSegments(in: requestContext.transcript)

if type != String.self {
let (jsonString, usage) = try await generateStructuredJSON(
session: session,
requestContext: requestContext,
prompt: prompt,
schema: schema,
options: options,
Expand All @@ -139,8 +140,8 @@
let tokens: [Int]
if let chatTemplateHandler = chatTemplateHandler {
// Use chat template handler with optional tools
let messages = chatTemplateHandler(session.instructions, prompt)
let toolSpecs: [ToolSpec]? = toolsHandler?(session.tools)
let messages = chatTemplateHandler(requestContext.instructions, prompt)
let toolSpecs: [ToolSpec]? = toolsHandler?(requestContext.tools)
tokens = try tokenizer.applyChatTemplate(messages: messages, tools: toolSpecs)
} else {
// Fall back to direct tokenizer encoding
Expand Down Expand Up @@ -227,17 +228,6 @@
}
}

// Validate that no image segments are present
do {
try validateNoImageSegments(in: session)
} catch {
return LanguageModelSession.ResponseStream(
stream: AsyncThrowingStream { continuation in
continuation.finish(throwing: error)
}
)
}

// Convert AnyLanguageModel GenerationOptions to swift-transformers GenerationConfig
let generationConfig = toGenerationConfig(options)

Expand All @@ -246,11 +236,13 @@
@Sendable continuation in
let task = Task {
do {
let requestContext = session.resolvedRequestContext()
try validateNoImageSegments(in: requestContext.transcript)
let tokens: [Int]
if let chatTemplateHandler = chatTemplateHandler {
// Use chat template handler with optional tools
let messages = chatTemplateHandler(session.instructions, prompt)
let toolSpecs: [ToolSpec]? = toolsHandler?(session.tools)
let messages = chatTemplateHandler(requestContext.instructions, prompt)
let toolSpecs: [ToolSpec]? = toolsHandler?(requestContext.tools)
tokens = try tokenizer.applyChatTemplate(messages: messages, tools: toolSpecs)
} else {
// Fall back to direct tokenizer encoding
Expand Down Expand Up @@ -298,10 +290,10 @@

// MARK: - Image Validation

private func validateNoImageSegments(in session: LanguageModelSession) throws {
private func validateNoImageSegments(in transcript: Transcript) throws {
// Note: Instructions is a plain text type without segments, so no image check needed there.
// Check for image segments in the most recent prompt
for entry in session.transcript.reversed() {
for entry in transcript.reversed() {
if case .prompt(let p) = entry {
for segment in p.segments {
if case .image = segment {
Expand Down Expand Up @@ -410,7 +402,7 @@
}

private func generateStructuredJSON(
session: LanguageModelSession,
requestContext: LanguageModelSession.RequestContext,
prompt: Prompt,
schema: GenerationSchema,
options: GenerationOptions,
Expand All @@ -420,7 +412,7 @@
var generationConfig = toStructuredGenerationConfig(options)

let promptTokens = try structuredPromptTokens(
in: session,
requestContext: requestContext,
prompt: prompt,
schema: schema,
includeSchemaInPrompt: includeSchemaInPrompt
Expand Down Expand Up @@ -457,20 +449,20 @@
}

private func structuredPromptTokens(
in session: LanguageModelSession,
requestContext: LanguageModelSession.RequestContext,
prompt: Prompt,
schema: GenerationSchema,
includeSchemaInPrompt: Bool
) throws -> [Int] {
if let chatTemplateHandler = chatTemplateHandler {
var messages = chatTemplateHandler(session.instructions, prompt)
var messages = chatTemplateHandler(requestContext.instructions, prompt)
if includeSchemaInPrompt {
let schemaPrompt = schemaPrompt(for: schema)
if !schemaPrompt.isEmpty {
messages.insert(["role": "system", "content": schemaPrompt], at: 0)
}
}
let toolSpecs: [ToolSpec]? = toolsHandler?(session.tools)
let toolSpecs: [ToolSpec]? = toolsHandler?(requestContext.tools)
return try tokenizer.applyChatTemplate(messages: messages, tools: toolSpecs)
}

Expand Down
30 changes: 7 additions & 23 deletions Sources/AnyLanguageModel/Models/FoundationLanguageModel.swift
Original file line number Diff line number Diff line change
Expand Up @@ -105,13 +105,13 @@
}

private func makeSession(
tools: [any FoundationModels.Tool],
transcript: FoundationModels.Transcript
for session: LanguageModelSession,
prompt: Prompt
) async throws -> FoundationModels.LanguageModelSession {
FoundationModels.LanguageModelSession(
makeFoundationModelsSession(
model: try await loadedModel(),
tools: tools,
transcript: transcript
session: session,
prompt: prompt
)
}

Expand Down Expand Up @@ -157,16 +157,8 @@
includeSchemaInPrompt: Bool,
options: GenerationOptions
) async throws -> LanguageModelSession.Response<Content> where Content: Generable {
let fmTools = session.tools.toFoundationModels()
let fmTranscript = fmTranscriptDroppingDuplicatePrompt(session.transcript, prompt: prompt)
.toFoundationModels(
instructions: session.instructions,
toolDefinitions: session.tools
.filter(\.includesSchemaInInstructions)
.map { Transcript.ToolDefinition(tool: $0) }
)
return try await fmRespond(
makeSession: { try await self.makeSession(tools: fmTools, transcript: fmTranscript) },
makeSession: { try await self.makeSession(for: session, prompt: prompt) },
fmPrompt: prompt.toFoundationModels(),
fmOptions: options.toFoundationModels(),
type: type,
Expand Down Expand Up @@ -217,16 +209,8 @@
includeSchemaInPrompt: Bool,
options: GenerationOptions
) -> sending LanguageModelSession.ResponseStream<Content> where Content: Generable {
let fmTools = session.tools.toFoundationModels()
let fmTranscript = fmTranscriptDroppingDuplicatePrompt(session.transcript, prompt: prompt)
.toFoundationModels(
instructions: session.instructions,
toolDefinitions: session.tools
.filter(\.includesSchemaInInstructions)
.map { Transcript.ToolDefinition(tool: $0) }
)
return fmStreamResponse(
makeSession: { try await self.makeSession(tools: fmTools, transcript: fmTranscript) },
makeSession: { try await self.makeSession(for: session, prompt: prompt) },
fmPrompt: prompt.toFoundationModels(),
fmOptions: options.toFoundationModels(),
type: type,
Expand Down
Loading
Loading