125 lines
4.3 KiB
Swift
125 lines
4.3 KiB
Swift
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
|
|
}
|
|
}
|