Files
music-assistant/Sources/MusicAssistantCore/Integrations/OpenAI/OpenAIClient.swift
T

164 lines
5.9 KiB
Swift

import Foundation
public struct OpenAIClientConfiguration: Equatable, Sendable {
public var apiKey: String
public var endpointURL: URL
public var model: String?
public var organizationID: String?
public var projectID: String?
public init(
apiKey: String,
endpointURL: URL,
model: String? = nil,
organizationID: String? = nil,
projectID: String? = nil
) {
self.apiKey = apiKey
self.endpointURL = endpointURL
self.model = model
self.organizationID = organizationID
self.projectID = projectID
}
}
public struct OpenAIClientRequest: Equatable, Sendable {
public var method: String
public var url: URL?
public var body: Data
public var additionalHeaders: [String: String]
public init(
method: String = "POST",
url: URL? = nil,
body: Data,
additionalHeaders: [String: String] = [:]
) {
self.method = method
self.url = url
self.body = body
self.additionalHeaders = additionalHeaders
}
}
public protocol OpenAIClientAdapter: Sendable {
func makeProjectGenerationRequest(
_ request: SongProjectGenerationRequest,
configuration: OpenAIClientConfiguration
) throws -> OpenAIClientRequest
func decodeProjectGenerationResult(from data: Data) throws -> SongProjectGenerationResult
func makeSongProjectDiscussionRequest(
_ request: SongProjectDiscussionRequest,
configuration: OpenAIClientConfiguration
) throws -> OpenAIClientRequest
func decodeSongProjectDiscussionResult(from data: Data) throws -> SongProjectDiscussionResult
func makeLyricsRevisionRequest(
_ request: LyricsRevisionRequest,
configuration: OpenAIClientConfiguration
) throws -> OpenAIClientRequest
func decodeLyricsRevisionResult(from data: Data) throws -> LyricsRevisionResult
func makeProjectUpdateRequest(
_ request: SongProjectUpdateRequest,
configuration: OpenAIClientConfiguration
) throws -> OpenAIClientRequest
func decodeProjectUpdateResult(from data: Data) throws -> SongProjectUpdateResult
}
public protocol OpenAIHTTPTransport: Sendable {
func data(for request: URLRequest) async throws -> (Data, HTTPURLResponse)
}
extension URLSession: OpenAIHTTPTransport {
public func data(for request: URLRequest) async throws -> (Data, HTTPURLResponse) {
let (data, response) = try await data(for: request, delegate: nil)
guard let httpResponse = response as? HTTPURLResponse else {
throw OpenAIClientError.invalidResponse
}
return (data, httpResponse)
}
}
public final class OpenAIClient: AIService, Sendable {
private let configuration: OpenAIClientConfiguration
private let adapter: any OpenAIClientAdapter
private let transport: any OpenAIHTTPTransport
public init(
configuration: OpenAIClientConfiguration,
adapter: any OpenAIClientAdapter,
transport: any OpenAIHTTPTransport = URLSession.shared
) {
self.configuration = configuration
self.adapter = adapter
self.transport = transport
}
public func generateSongProject(from request: SongProjectGenerationRequest) async throws -> SongProjectGenerationResult {
let clientRequest = try adapter.makeProjectGenerationRequest(request, configuration: configuration)
let data = try await perform(clientRequest)
return try adapter.decodeProjectGenerationResult(from: data)
}
public func discussSongProject(from request: SongProjectDiscussionRequest) async throws -> SongProjectDiscussionResult {
let clientRequest = try adapter.makeSongProjectDiscussionRequest(request, configuration: configuration)
let data = try await perform(clientRequest)
return try adapter.decodeSongProjectDiscussionResult(from: data)
}
public func reviseLyrics(from request: LyricsRevisionRequest) async throws -> LyricsRevisionResult {
let clientRequest = try adapter.makeLyricsRevisionRequest(request, configuration: configuration)
let data = try await perform(clientRequest)
return try adapter.decodeLyricsRevisionResult(from: data)
}
public func proposeProjectUpdate(from request: SongProjectUpdateRequest) async throws -> SongProjectUpdateResult {
let clientRequest = try adapter.makeProjectUpdateRequest(request, configuration: configuration)
let data = try await perform(clientRequest)
return try adapter.decodeProjectUpdateResult(from: data)
}
private func perform(_ clientRequest: OpenAIClientRequest) async throws -> Data {
guard !configuration.apiKey.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty else {
throw OpenAIClientError.missingAPIKey
}
var request = URLRequest(url: clientRequest.url ?? configuration.endpointURL)
request.httpMethod = clientRequest.method
request.httpBody = clientRequest.body
request.setValue("Bearer \(configuration.apiKey)", forHTTPHeaderField: "Authorization")
request.setValue("application/json", forHTTPHeaderField: "Content-Type")
if let organizationID = configuration.organizationID {
request.setValue(organizationID, forHTTPHeaderField: "OpenAI-Organization")
}
if let projectID = configuration.projectID {
request.setValue(projectID, forHTTPHeaderField: "OpenAI-Project")
}
for (header, value) in clientRequest.additionalHeaders {
request.setValue(value, forHTTPHeaderField: header)
}
let (data, response) = try await transport.data(for: request)
guard (200..<300).contains(response.statusCode) else {
throw OpenAIClientError.unacceptableStatusCode(response.statusCode)
}
return data
}
}
public enum OpenAIClientError: Error, Equatable, Sendable {
case missingAPIKey
case invalidResponse
case unacceptableStatusCode(Int)
}