diff --git a/Sources/MusicAssistantCore/Integrations/OpenAI/OpenAIClient.swift b/Sources/MusicAssistantCore/Integrations/OpenAI/OpenAIClient.swift index 586047d..8fbf74a 100644 --- a/Sources/MusicAssistantCore/Integrations/OpenAI/OpenAIClient.swift +++ b/Sources/MusicAssistantCore/Integrations/OpenAI/OpenAIClient.swift @@ -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) } diff --git a/Sources/MusicAssistantCore/Services/AI/AIServiceRetrying.swift b/Sources/MusicAssistantCore/Services/AI/AIServiceRetrying.swift new file mode 100644 index 0000000..91ad256 --- /dev/null +++ b/Sources/MusicAssistantCore/Services/AI/AIServiceRetrying.swift @@ -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( + _ 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 + } +} diff --git a/Tests/MusicAssistantCoreTests/OpenAIClientTests.swift b/Tests/MusicAssistantCoreTests/OpenAIClientTests.swift index f3e3a4f..ab29a1c 100644 --- a/Tests/MusicAssistantCoreTests/OpenAIClientTests.swift +++ b/Tests/MusicAssistantCoreTests/OpenAIClientTests.swift @@ -128,7 +128,7 @@ final class OpenAIClientTests: XCTestCase { XCTAssertEqual(recordedRequestCount, 0) } - func testClientReportsUnacceptableStatusWithoutLeakingResponseBody() async throws { + func testClientReportsRateLimitWithoutLeakingResponseBody() async throws { let transport = RecordingOpenAITransport( statusCode: 429, responseData: Data("secret server detail".utf8) @@ -150,7 +150,32 @@ final class OpenAIClientTests: XCTestCase { ) XCTFail("Expected status error to throw.") } 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 let statusCode: Int 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.responseData = responseData + self.responseHeaders = responseHeaders } var requestBodyString: String? { @@ -176,7 +207,7 @@ private actor RecordingOpenAITransport: OpenAIHTTPTransport { url: request.url!, statusCode: statusCode, httpVersion: "HTTP/1.1", - headerFields: nil + headerFields: responseHeaders )! return (responseData, response) } diff --git a/Tests/MusicAssistantCoreTests/RetryingAIServiceTests.swift b/Tests/MusicAssistantCoreTests/RetryingAIServiceTests.swift new file mode 100644 index 0000000..de1c4b9 --- /dev/null +++ b/Tests/MusicAssistantCoreTests/RetryingAIServiceTests.swift @@ -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() + } + } +} diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index a779303..d364f4a 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -49,6 +49,9 @@ questions without applying a project update. 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 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 diff --git a/docs/TASKS.md b/docs/TASKS.md index 46da6e2..958d776 100644 --- a/docs/TASKS.md +++ b/docs/TASKS.md @@ -63,7 +63,7 @@ requirement is missing and blocks implementation, record it in and production decisions. - [x] Implement optional Discuss mode. - [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