Add resilient AI service handling
This commit is contained in:
@@ -125,6 +125,8 @@ public final class OpenAIClient: AIService, Sendable {
|
||||
}
|
||||
|
||||
private func perform(_ clientRequest: OpenAIClientRequest) async throws -> Data {
|
||||
try Task.checkCancellation()
|
||||
|
||||
guard !configuration.apiKey.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty else {
|
||||
throw OpenAIClientError.missingAPIKey
|
||||
}
|
||||
@@ -149,15 +151,30 @@ public final class OpenAIClient: AIService, Sendable {
|
||||
|
||||
let (data, response) = try await transport.data(for: request)
|
||||
guard (200..<300).contains(response.statusCode) else {
|
||||
if response.statusCode == 429 {
|
||||
throw OpenAIClientError.rateLimited(
|
||||
retryAfter: Self.retryAfterInterval(from: response)
|
||||
)
|
||||
}
|
||||
if (500..<600).contains(response.statusCode) {
|
||||
throw OpenAIClientError.serverError(response.statusCode)
|
||||
}
|
||||
throw OpenAIClientError.unacceptableStatusCode(response.statusCode)
|
||||
}
|
||||
|
||||
return data
|
||||
}
|
||||
|
||||
private static func retryAfterInterval(from response: HTTPURLResponse) -> TimeInterval? {
|
||||
guard let value = response.value(forHTTPHeaderField: "Retry-After") else { return nil }
|
||||
return TimeInterval(value.trimmingCharacters(in: .whitespacesAndNewlines))
|
||||
}
|
||||
}
|
||||
|
||||
public enum OpenAIClientError: Error, Equatable, Sendable {
|
||||
case missingAPIKey
|
||||
case invalidResponse
|
||||
case rateLimited(retryAfter: TimeInterval?)
|
||||
case serverError(Int)
|
||||
case unacceptableStatusCode(Int)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
import Foundation
|
||||
|
||||
public struct AIServiceRetryPolicy: Equatable, Sendable {
|
||||
public static let `default` = AIServiceRetryPolicy()
|
||||
|
||||
public var maxAttempts: Int
|
||||
public var baseDelay: TimeInterval
|
||||
public var maximumDelay: TimeInterval
|
||||
|
||||
public init(
|
||||
maxAttempts: Int = 3,
|
||||
baseDelay: TimeInterval = 0.5,
|
||||
maximumDelay: TimeInterval = 8
|
||||
) {
|
||||
self.maxAttempts = max(1, maxAttempts)
|
||||
self.baseDelay = max(0, baseDelay)
|
||||
self.maximumDelay = max(0, maximumDelay)
|
||||
}
|
||||
|
||||
func delay(forFailedAttempt attempt: Int, retryAfter: TimeInterval? = nil) -> TimeInterval {
|
||||
if let retryAfter {
|
||||
return max(0, retryAfter)
|
||||
}
|
||||
|
||||
let multiplier = pow(2, Double(max(0, attempt - 1)))
|
||||
return min(maximumDelay, baseDelay * multiplier)
|
||||
}
|
||||
}
|
||||
|
||||
public protocol AIServiceRetryScheduler: Sendable {
|
||||
func sleep(for interval: TimeInterval) async throws
|
||||
}
|
||||
|
||||
public struct TaskAIServiceRetryScheduler: AIServiceRetryScheduler {
|
||||
public init() {}
|
||||
|
||||
public func sleep(for interval: TimeInterval) async throws {
|
||||
guard interval > 0 else { return }
|
||||
let maximumInterval = TimeInterval(UInt64.max / 1_000_000_000)
|
||||
let clampedInterval = min(interval, maximumInterval)
|
||||
let nanoseconds = UInt64((clampedInterval * 1_000_000_000).rounded())
|
||||
try await Task.sleep(nanoseconds: nanoseconds)
|
||||
}
|
||||
}
|
||||
|
||||
public final class RetryingAIService: AIService, Sendable {
|
||||
private let baseService: any AIService
|
||||
private let policy: AIServiceRetryPolicy
|
||||
private let scheduler: any AIServiceRetryScheduler
|
||||
|
||||
public init(
|
||||
baseService: any AIService,
|
||||
policy: AIServiceRetryPolicy = .default,
|
||||
scheduler: any AIServiceRetryScheduler = TaskAIServiceRetryScheduler()
|
||||
) {
|
||||
self.baseService = baseService
|
||||
self.policy = policy
|
||||
self.scheduler = scheduler
|
||||
}
|
||||
|
||||
public func generateSongProject(from request: SongProjectGenerationRequest) async throws -> SongProjectGenerationResult {
|
||||
try await perform { try await self.baseService.generateSongProject(from: request) }
|
||||
}
|
||||
|
||||
public func discussSongProject(from request: SongProjectDiscussionRequest) async throws -> SongProjectDiscussionResult {
|
||||
try await perform { try await self.baseService.discussSongProject(from: request) }
|
||||
}
|
||||
|
||||
public func reviseLyrics(from request: LyricsRevisionRequest) async throws -> LyricsRevisionResult {
|
||||
try await perform { try await self.baseService.reviseLyrics(from: request) }
|
||||
}
|
||||
|
||||
public func proposeProjectUpdate(from request: SongProjectUpdateRequest) async throws -> SongProjectUpdateResult {
|
||||
try await perform { try await self.baseService.proposeProjectUpdate(from: request) }
|
||||
}
|
||||
|
||||
private func perform<Value: Sendable>(
|
||||
_ operation: @escaping @Sendable () async throws -> Value
|
||||
) async throws -> Value {
|
||||
var attempt = 0
|
||||
|
||||
while true {
|
||||
try Task.checkCancellation()
|
||||
attempt += 1
|
||||
|
||||
do {
|
||||
return try await operation()
|
||||
} catch is CancellationError {
|
||||
throw CancellationError()
|
||||
} catch {
|
||||
guard attempt < policy.maxAttempts,
|
||||
let delay = retryDelay(for: error, failedAttempt: attempt) else {
|
||||
throw error
|
||||
}
|
||||
|
||||
try await scheduler.sleep(for: delay)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private func retryDelay(for error: Error, failedAttempt attempt: Int) -> TimeInterval? {
|
||||
if let error = error as? OpenAIClientError {
|
||||
switch error {
|
||||
case let .rateLimited(retryAfter):
|
||||
return policy.delay(forFailedAttempt: attempt, retryAfter: retryAfter)
|
||||
case .serverError:
|
||||
return policy.delay(forFailedAttempt: attempt)
|
||||
case .missingAPIKey, .invalidResponse, .unacceptableStatusCode:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
if let error = error as? URLError {
|
||||
switch error.code {
|
||||
case .timedOut, .networkConnectionLost, .notConnectedToInternet, .cannotConnectToHost, .cannotFindHost, .dnsLookupFailed:
|
||||
return policy.delay(forFailedAttempt: attempt)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user