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() } } }