151 lines
5.3 KiB
Swift
151 lines
5.3 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 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 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)
|
|
}
|