Add resilient AI service handling

This commit is contained in:
diyaa
2026-09-13 21:28:13 +02:00
parent ae88c6e533
commit eb61bd076c
6 changed files with 324 additions and 5 deletions
@@ -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)
}
@@ -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()
}
}
}