Add resilient AI service handling
This commit is contained in:
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user