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 {
|
private func perform(_ clientRequest: OpenAIClientRequest) async throws -> Data {
|
||||||
|
try Task.checkCancellation()
|
||||||
|
|
||||||
guard !configuration.apiKey.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty else {
|
guard !configuration.apiKey.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty else {
|
||||||
throw OpenAIClientError.missingAPIKey
|
throw OpenAIClientError.missingAPIKey
|
||||||
}
|
}
|
||||||
@@ -149,15 +151,30 @@ public final class OpenAIClient: AIService, Sendable {
|
|||||||
|
|
||||||
let (data, response) = try await transport.data(for: request)
|
let (data, response) = try await transport.data(for: request)
|
||||||
guard (200..<300).contains(response.statusCode) else {
|
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)
|
throw OpenAIClientError.unacceptableStatusCode(response.statusCode)
|
||||||
}
|
}
|
||||||
|
|
||||||
return data
|
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 {
|
public enum OpenAIClientError: Error, Equatable, Sendable {
|
||||||
case missingAPIKey
|
case missingAPIKey
|
||||||
case invalidResponse
|
case invalidResponse
|
||||||
|
case rateLimited(retryAfter: TimeInterval?)
|
||||||
|
case serverError(Int)
|
||||||
case unacceptableStatusCode(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
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -128,7 +128,7 @@ final class OpenAIClientTests: XCTestCase {
|
|||||||
XCTAssertEqual(recordedRequestCount, 0)
|
XCTAssertEqual(recordedRequestCount, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
func testClientReportsUnacceptableStatusWithoutLeakingResponseBody() async throws {
|
func testClientReportsRateLimitWithoutLeakingResponseBody() async throws {
|
||||||
let transport = RecordingOpenAITransport(
|
let transport = RecordingOpenAITransport(
|
||||||
statusCode: 429,
|
statusCode: 429,
|
||||||
responseData: Data("secret server detail".utf8)
|
responseData: Data("secret server detail".utf8)
|
||||||
@@ -150,7 +150,32 @@ final class OpenAIClientTests: XCTestCase {
|
|||||||
)
|
)
|
||||||
XCTFail("Expected status error to throw.")
|
XCTFail("Expected status error to throw.")
|
||||||
} catch let error as OpenAIClientError {
|
} catch let error as OpenAIClientError {
|
||||||
XCTAssertEqual(error, .unacceptableStatusCode(429))
|
XCTAssertEqual(error, .rateLimited(retryAfter: nil))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func testClientReadsRetryAfterForRateLimitedResponses() async throws {
|
||||||
|
let transport = RecordingOpenAITransport(
|
||||||
|
statusCode: 429,
|
||||||
|
responseData: Data(),
|
||||||
|
responseHeaders: ["Retry-After": "3.5"]
|
||||||
|
)
|
||||||
|
let client = OpenAIClient(
|
||||||
|
configuration: OpenAIClientConfiguration(
|
||||||
|
apiKey: "test-api-key",
|
||||||
|
endpointURL: URL(string: "https://api.example.test/v1/configured")!
|
||||||
|
),
|
||||||
|
adapter: MockOpenAIClientAdapter(),
|
||||||
|
transport: transport
|
||||||
|
)
|
||||||
|
|
||||||
|
do {
|
||||||
|
_ = try await client.generateSongProject(
|
||||||
|
from: SongProjectGenerationRequest(context: AIRequestContext(userInstruction: "Generate."))
|
||||||
|
)
|
||||||
|
XCTFail("Expected rate limit error to throw.")
|
||||||
|
} catch let error as OpenAIClientError {
|
||||||
|
XCTAssertEqual(error, .rateLimited(retryAfter: 3.5))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -159,10 +184,16 @@ private actor RecordingOpenAITransport: OpenAIHTTPTransport {
|
|||||||
private(set) var recordedRequests: [URLRequest] = []
|
private(set) var recordedRequests: [URLRequest] = []
|
||||||
private let statusCode: Int
|
private let statusCode: Int
|
||||||
private let responseData: Data
|
private let responseData: Data
|
||||||
|
private let responseHeaders: [String: String]?
|
||||||
|
|
||||||
init(statusCode: Int, responseData: Data) {
|
init(
|
||||||
|
statusCode: Int,
|
||||||
|
responseData: Data,
|
||||||
|
responseHeaders: [String: String]? = nil
|
||||||
|
) {
|
||||||
self.statusCode = statusCode
|
self.statusCode = statusCode
|
||||||
self.responseData = responseData
|
self.responseData = responseData
|
||||||
|
self.responseHeaders = responseHeaders
|
||||||
}
|
}
|
||||||
|
|
||||||
var requestBodyString: String? {
|
var requestBodyString: String? {
|
||||||
@@ -176,7 +207,7 @@ private actor RecordingOpenAITransport: OpenAIHTTPTransport {
|
|||||||
url: request.url!,
|
url: request.url!,
|
||||||
statusCode: statusCode,
|
statusCode: statusCode,
|
||||||
httpVersion: "HTTP/1.1",
|
httpVersion: "HTTP/1.1",
|
||||||
headerFields: nil
|
headerFields: responseHeaders
|
||||||
)!
|
)!
|
||||||
return (responseData, response)
|
return (responseData, response)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,144 @@
|
|||||||
|
import MusicAssistantCore
|
||||||
|
import XCTest
|
||||||
|
|
||||||
|
final class RetryingAIServiceTests: XCTestCase {
|
||||||
|
func testRetriesTemporaryServerFailureThenReturnsResult() async throws {
|
||||||
|
let baseService = ScriptedRetryAIService(
|
||||||
|
outcomes: [.serverFailure(503), .serverFailure(503), .success("Generated")]
|
||||||
|
)
|
||||||
|
let scheduler = RecordingRetryScheduler()
|
||||||
|
let service = RetryingAIService(
|
||||||
|
baseService: baseService,
|
||||||
|
policy: AIServiceRetryPolicy(maxAttempts: 3, baseDelay: 0.25, maximumDelay: 2),
|
||||||
|
scheduler: scheduler
|
||||||
|
)
|
||||||
|
|
||||||
|
let result = try await service.generateSongProject(
|
||||||
|
from: SongProjectGenerationRequest(context: AIRequestContext(userInstruction: "Generate."))
|
||||||
|
)
|
||||||
|
let attemptCount = await baseService.generateAttemptCount
|
||||||
|
let delays = await scheduler.delays
|
||||||
|
|
||||||
|
XCTAssertEqual(result.project.title, "Generated")
|
||||||
|
XCTAssertEqual(attemptCount, 3)
|
||||||
|
XCTAssertEqual(delays, [0.25, 0.5])
|
||||||
|
}
|
||||||
|
|
||||||
|
func testHonorsProviderRateLimitDelayBeforeRetrying() async throws {
|
||||||
|
let baseService = ScriptedRetryAIService(
|
||||||
|
outcomes: [.rateLimited(4), .success("Generated")]
|
||||||
|
)
|
||||||
|
let scheduler = RecordingRetryScheduler()
|
||||||
|
let service = RetryingAIService(
|
||||||
|
baseService: baseService,
|
||||||
|
policy: AIServiceRetryPolicy(maxAttempts: 3, baseDelay: 0.25, maximumDelay: 2),
|
||||||
|
scheduler: scheduler
|
||||||
|
)
|
||||||
|
|
||||||
|
_ = try await service.generateSongProject(
|
||||||
|
from: SongProjectGenerationRequest(context: AIRequestContext(userInstruction: "Generate."))
|
||||||
|
)
|
||||||
|
let delays = await scheduler.delays
|
||||||
|
|
||||||
|
XCTAssertEqual(delays, [4])
|
||||||
|
}
|
||||||
|
|
||||||
|
func testDoesNotRetryPermanentFailures() async {
|
||||||
|
let baseService = ScriptedRetryAIService(outcomes: [.missingAPIKey])
|
||||||
|
let scheduler = RecordingRetryScheduler()
|
||||||
|
let service = RetryingAIService(baseService: baseService, scheduler: scheduler)
|
||||||
|
|
||||||
|
do {
|
||||||
|
_ = try await service.generateSongProject(
|
||||||
|
from: SongProjectGenerationRequest(context: AIRequestContext(userInstruction: "Generate."))
|
||||||
|
)
|
||||||
|
XCTFail("Expected missing key error to throw.")
|
||||||
|
} catch let error as OpenAIClientError {
|
||||||
|
XCTAssertEqual(error, .missingAPIKey)
|
||||||
|
} catch {
|
||||||
|
XCTFail("Expected OpenAIClientError.")
|
||||||
|
}
|
||||||
|
|
||||||
|
let attemptCount = await baseService.generateAttemptCount
|
||||||
|
let delays = await scheduler.delays
|
||||||
|
XCTAssertEqual(attemptCount, 1)
|
||||||
|
XCTAssertTrue(delays.isEmpty)
|
||||||
|
}
|
||||||
|
|
||||||
|
func testCancellationWhileWaitingStopsFurtherRetries() async {
|
||||||
|
let baseService = ScriptedRetryAIService(outcomes: [.serverFailure(503), .success("Unused")])
|
||||||
|
let scheduler = RecordingRetryScheduler(shouldCancel: true)
|
||||||
|
let service = RetryingAIService(baseService: baseService, scheduler: scheduler)
|
||||||
|
|
||||||
|
do {
|
||||||
|
_ = try await service.generateSongProject(
|
||||||
|
from: SongProjectGenerationRequest(context: AIRequestContext(userInstruction: "Generate."))
|
||||||
|
)
|
||||||
|
XCTFail("Expected cancellation to throw.")
|
||||||
|
} catch is CancellationError {
|
||||||
|
} catch {
|
||||||
|
XCTFail("Expected CancellationError.")
|
||||||
|
}
|
||||||
|
|
||||||
|
let attemptCount = await baseService.generateAttemptCount
|
||||||
|
XCTAssertEqual(attemptCount, 1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private enum RetryOutcome: Sendable {
|
||||||
|
case serverFailure(Int)
|
||||||
|
case rateLimited(TimeInterval?)
|
||||||
|
case missingAPIKey
|
||||||
|
case success(String)
|
||||||
|
}
|
||||||
|
|
||||||
|
private actor ScriptedRetryAIService: AIService {
|
||||||
|
private var outcomes: [RetryOutcome]
|
||||||
|
private(set) var generateAttemptCount = 0
|
||||||
|
|
||||||
|
init(outcomes: [RetryOutcome]) {
|
||||||
|
self.outcomes = outcomes
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateSongProject(from request: SongProjectGenerationRequest) async throws -> SongProjectGenerationResult {
|
||||||
|
generateAttemptCount += 1
|
||||||
|
guard !outcomes.isEmpty else {
|
||||||
|
return SongProjectGenerationResult(project: SongProject(title: "Unexpected", idea: request.context.userInstruction))
|
||||||
|
}
|
||||||
|
|
||||||
|
switch outcomes.removeFirst() {
|
||||||
|
case let .serverFailure(statusCode):
|
||||||
|
throw OpenAIClientError.serverError(statusCode)
|
||||||
|
case let .rateLimited(retryAfter):
|
||||||
|
throw OpenAIClientError.rateLimited(retryAfter: retryAfter)
|
||||||
|
case .missingAPIKey:
|
||||||
|
throw OpenAIClientError.missingAPIKey
|
||||||
|
case let .success(title):
|
||||||
|
return SongProjectGenerationResult(project: SongProject(title: title, idea: request.context.userInstruction))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func reviseLyrics(from request: LyricsRevisionRequest) async throws -> LyricsRevisionResult {
|
||||||
|
LyricsRevisionResult(lyrics: request.sourceLyrics)
|
||||||
|
}
|
||||||
|
|
||||||
|
func proposeProjectUpdate(from request: SongProjectUpdateRequest) async throws -> SongProjectUpdateResult {
|
||||||
|
SongProjectUpdateResult(project: request.project)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private actor RecordingRetryScheduler: AIServiceRetryScheduler {
|
||||||
|
private(set) var delays: [TimeInterval] = []
|
||||||
|
private let shouldCancel: Bool
|
||||||
|
|
||||||
|
init(shouldCancel: Bool = false) {
|
||||||
|
self.shouldCancel = shouldCancel
|
||||||
|
}
|
||||||
|
|
||||||
|
func sleep(for interval: TimeInterval) async throws {
|
||||||
|
delays.append(interval)
|
||||||
|
if shouldCancel {
|
||||||
|
throw CancellationError()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -49,6 +49,9 @@ questions without applying a project update.
|
|||||||
All structured AI updates pass through a scope-aware merger. It applies only
|
All structured AI updates pass through a scope-aware merger. It applies only
|
||||||
the scopes declared by the response, ignores user-locked scopes and preserves
|
the scopes declared by the response, ignores user-locked scopes and preserves
|
||||||
manual control values.
|
manual control values.
|
||||||
|
AI provider calls can be wrapped by a configurable retry service. It retries only
|
||||||
|
temporary network, server and rate-limit failures, honors a provider-supplied
|
||||||
|
retry delay when available, and propagates cancellation without retrying.
|
||||||
|
|
||||||
## Prompt Compiler
|
## Prompt Compiler
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -63,7 +63,7 @@ requirement is missing and blocks implementation, record it in
|
|||||||
and production decisions.
|
and production decisions.
|
||||||
- [x] Implement optional Discuss mode.
|
- [x] Implement optional Discuss mode.
|
||||||
- [x] Enforce user-lock/manual-value precedence over AI output.
|
- [x] Enforce user-lock/manual-value precedence over AI output.
|
||||||
- [ ] Add error, retry, cancellation and rate-limit handling.
|
- [x] Add error, retry, cancellation and rate-limit handling.
|
||||||
|
|
||||||
## Phase 5 --- Arabic Lyrics Processing
|
## Phase 5 --- Arabic Lyrics Processing
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user