288 lines
11 KiB
Swift
288 lines
11 KiB
Swift
import Foundation
|
|
import MusicAssistantCore
|
|
import XCTest
|
|
|
|
final class OpenAIClientTests: XCTestCase {
|
|
func testGenerateSongProjectUsesConfiguredEndpointAndHeaders() async throws {
|
|
let transport = RecordingOpenAITransport(
|
|
statusCode: 200,
|
|
responseData: try JSONEncoder().encode(MockOpenAIResponse(text: "Generated"))
|
|
)
|
|
let configuration = OpenAIClientConfiguration(
|
|
apiKey: "test-api-key",
|
|
endpointURL: URL(string: "https://api.example.test/v1/configured")!,
|
|
model: "configured-model",
|
|
organizationID: "org-test",
|
|
projectID: "project-test"
|
|
)
|
|
let client = OpenAIClient(
|
|
configuration: configuration,
|
|
adapter: MockOpenAIClientAdapter(),
|
|
transport: transport
|
|
)
|
|
|
|
let result = try await client.generateSongProject(
|
|
from: SongProjectGenerationRequest(
|
|
context: AIRequestContext(userInstruction: "Generate this."),
|
|
discussionMode: .auto
|
|
)
|
|
)
|
|
let request = await transport.recordedRequests.first
|
|
|
|
XCTAssertEqual(result.project.title, "Generated")
|
|
XCTAssertEqual(request?.url, configuration.endpointURL)
|
|
XCTAssertEqual(request?.httpMethod, "POST")
|
|
XCTAssertEqual(request?.value(forHTTPHeaderField: "Authorization"), "Bearer test-api-key")
|
|
XCTAssertEqual(request?.value(forHTTPHeaderField: "OpenAI-Organization"), "org-test")
|
|
XCTAssertEqual(request?.value(forHTTPHeaderField: "OpenAI-Project"), "project-test")
|
|
XCTAssertEqual(request?.value(forHTTPHeaderField: "X-Adapter"), "mock")
|
|
XCTAssertEqual(request?.value(forHTTPHeaderField: "Content-Type"), "application/json")
|
|
let requestBodyString = await transport.requestBodyString
|
|
XCTAssertEqual(requestBodyString, "generate|configured-model|Generate this.")
|
|
}
|
|
|
|
func testClientUsesAdapterURLOverrideWhenProvided() async throws {
|
|
let overrideURL = URL(string: "https://api.example.test/v1/override")!
|
|
let transport = RecordingOpenAITransport(
|
|
statusCode: 200,
|
|
responseData: try JSONEncoder().encode(MockOpenAIResponse(text: "Updated"))
|
|
)
|
|
let client = OpenAIClient(
|
|
configuration: OpenAIClientConfiguration(
|
|
apiKey: "test-api-key",
|
|
endpointURL: URL(string: "https://api.example.test/v1/configured")!
|
|
),
|
|
adapter: MockOpenAIClientAdapter(overrideURL: overrideURL),
|
|
transport: transport
|
|
)
|
|
|
|
_ = try await client.proposeProjectUpdate(
|
|
from: SongProjectUpdateRequest(
|
|
context: AIRequestContext(userInstruction: "Update it."),
|
|
project: SongProject(title: "Original", idea: "Original idea"),
|
|
allowedScopes: [.lyrics]
|
|
)
|
|
)
|
|
|
|
let recordedURL = await transport.recordedRequests.first?.url
|
|
XCTAssertEqual(recordedURL, overrideURL)
|
|
}
|
|
|
|
func testDiscussSongProjectUsesAdapterAndReturnsQuestions() async throws {
|
|
let transport = RecordingOpenAITransport(
|
|
statusCode: 200,
|
|
responseData: try JSONEncoder().encode(MockOpenAIResponse(text: "Should the chorus be intimate or anthemic?"))
|
|
)
|
|
let configuration = OpenAIClientConfiguration(
|
|
apiKey: "test-api-key",
|
|
endpointURL: URL(string: "https://api.example.test/v1/configured")!,
|
|
model: "configured-model"
|
|
)
|
|
let client = OpenAIClient(
|
|
configuration: configuration,
|
|
adapter: MockOpenAIClientAdapter(),
|
|
transport: transport
|
|
)
|
|
|
|
let result = try await client.discussSongProject(
|
|
from: SongProjectDiscussionRequest(
|
|
context: AIRequestContext(userInstruction: "Ask before deciding."),
|
|
project: SongProject(title: "Discussion", idea: "Plan this", conversationMode: .discuss)
|
|
)
|
|
)
|
|
let requestBodyString = await transport.requestBodyString
|
|
|
|
XCTAssertEqual(result.questions, ["Should the chorus be intimate or anthemic?"])
|
|
XCTAssertEqual(requestBodyString, "discuss|configured-model|Ask before deciding.")
|
|
}
|
|
|
|
func testClientRejectsMissingAPIKeyBeforeSendingRequest() async throws {
|
|
let transport = RecordingOpenAITransport(
|
|
statusCode: 200,
|
|
responseData: try JSONEncoder().encode(MockOpenAIResponse(text: "Never sent"))
|
|
)
|
|
let client = OpenAIClient(
|
|
configuration: OpenAIClientConfiguration(
|
|
apiKey: " ",
|
|
endpointURL: URL(string: "https://api.example.test/v1/configured")!
|
|
),
|
|
adapter: MockOpenAIClientAdapter(),
|
|
transport: transport
|
|
)
|
|
|
|
do {
|
|
_ = try await client.reviseLyrics(
|
|
from: LyricsRevisionRequest(
|
|
context: AIRequestContext(userInstruction: "Improve."),
|
|
project: SongProject(title: "Song", idea: "Idea"),
|
|
sourceLyrics: Lyrics(text: "Draft"),
|
|
mode: .improve
|
|
)
|
|
)
|
|
XCTFail("Expected missing API key to throw.")
|
|
} catch let error as OpenAIClientError {
|
|
XCTAssertEqual(error, .missingAPIKey)
|
|
}
|
|
|
|
let recordedRequestCount = await transport.recordedRequests.count
|
|
XCTAssertEqual(recordedRequestCount, 0)
|
|
}
|
|
|
|
func testClientReportsRateLimitWithoutLeakingResponseBody() async throws {
|
|
let transport = RecordingOpenAITransport(
|
|
statusCode: 429,
|
|
responseData: Data("secret server detail".utf8)
|
|
)
|
|
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 status error to throw.")
|
|
} catch let error as OpenAIClientError {
|
|
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))
|
|
}
|
|
}
|
|
}
|
|
|
|
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,
|
|
responseHeaders: [String: String]? = nil
|
|
) {
|
|
self.statusCode = statusCode
|
|
self.responseData = responseData
|
|
self.responseHeaders = responseHeaders
|
|
}
|
|
|
|
var requestBodyString: String? {
|
|
guard let body = recordedRequests.first?.httpBody else { return nil }
|
|
return String(data: body, encoding: .utf8)
|
|
}
|
|
|
|
func data(for request: URLRequest) async throws -> (Data, HTTPURLResponse) {
|
|
recordedRequests.append(request)
|
|
let response = HTTPURLResponse(
|
|
url: request.url!,
|
|
statusCode: statusCode,
|
|
httpVersion: "HTTP/1.1",
|
|
headerFields: responseHeaders
|
|
)!
|
|
return (responseData, response)
|
|
}
|
|
}
|
|
|
|
private struct MockOpenAIClientAdapter: OpenAIClientAdapter {
|
|
var overrideURL: URL?
|
|
|
|
func makeProjectGenerationRequest(
|
|
_ request: SongProjectGenerationRequest,
|
|
configuration: OpenAIClientConfiguration
|
|
) throws -> OpenAIClientRequest {
|
|
OpenAIClientRequest(
|
|
url: overrideURL,
|
|
body: Data("generate|\(configuration.model ?? "")|\(request.context.userInstruction)".utf8),
|
|
additionalHeaders: ["X-Adapter": "mock"]
|
|
)
|
|
}
|
|
|
|
func decodeProjectGenerationResult(from data: Data) throws -> SongProjectGenerationResult {
|
|
let response = try JSONDecoder().decode(MockOpenAIResponse.self, from: data)
|
|
return SongProjectGenerationResult(
|
|
project: SongProject(title: response.text, idea: response.text)
|
|
)
|
|
}
|
|
|
|
func makeSongProjectDiscussionRequest(
|
|
_ request: SongProjectDiscussionRequest,
|
|
configuration: OpenAIClientConfiguration
|
|
) throws -> OpenAIClientRequest {
|
|
OpenAIClientRequest(
|
|
url: overrideURL,
|
|
body: Data("discuss|\(configuration.model ?? "")|\(request.context.userInstruction)".utf8)
|
|
)
|
|
}
|
|
|
|
func decodeSongProjectDiscussionResult(from data: Data) throws -> SongProjectDiscussionResult {
|
|
let response = try JSONDecoder().decode(MockOpenAIResponse.self, from: data)
|
|
return SongProjectDiscussionResult(questions: [response.text])
|
|
}
|
|
|
|
func makeLyricsRevisionRequest(
|
|
_ request: LyricsRevisionRequest,
|
|
configuration: OpenAIClientConfiguration
|
|
) throws -> OpenAIClientRequest {
|
|
OpenAIClientRequest(
|
|
url: overrideURL,
|
|
body: Data("lyrics|\(configuration.model ?? "")|\(request.context.userInstruction)".utf8)
|
|
)
|
|
}
|
|
|
|
func decodeLyricsRevisionResult(from data: Data) throws -> LyricsRevisionResult {
|
|
let response = try JSONDecoder().decode(MockOpenAIResponse.self, from: data)
|
|
return LyricsRevisionResult(lyrics: Lyrics(text: response.text))
|
|
}
|
|
|
|
func makeProjectUpdateRequest(
|
|
_ request: SongProjectUpdateRequest,
|
|
configuration: OpenAIClientConfiguration
|
|
) throws -> OpenAIClientRequest {
|
|
OpenAIClientRequest(
|
|
url: overrideURL,
|
|
body: Data("update|\(configuration.model ?? "")|\(request.context.userInstruction)".utf8)
|
|
)
|
|
}
|
|
|
|
func decodeProjectUpdateResult(from data: Data) throws -> SongProjectUpdateResult {
|
|
let response = try JSONDecoder().decode(MockOpenAIResponse.self, from: data)
|
|
return SongProjectUpdateResult(
|
|
project: SongProject(title: response.text, idea: response.text)
|
|
)
|
|
}
|
|
}
|
|
|
|
private struct MockOpenAIResponse: Codable {
|
|
var text: String
|
|
}
|