|
@@ -21,14 +21,29 @@ enum OpenAIServiceError: LocalizedError {
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
protocol OpenAIChatServicing {
|
|
protocol OpenAIChatServicing {
|
|
|
- func completeText(systemPrompt: String, userPrompt: String) async throws -> String
|
|
|
|
|
|
|
+ func completeText(systemPrompt: String, userPrompt: String, maxTokens: Int?) async throws -> String
|
|
|
func completeJSON<T: Decodable>(
|
|
func completeJSON<T: Decodable>(
|
|
|
systemPrompt: String,
|
|
systemPrompt: String,
|
|
|
userPrompt: String,
|
|
userPrompt: String,
|
|
|
|
|
+ maxTokens: Int?,
|
|
|
as type: T.Type
|
|
as type: T.Type
|
|
|
) async throws -> T
|
|
) async throws -> T
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
|
|
+extension OpenAIChatServicing {
|
|
|
|
|
+ func completeText(systemPrompt: String, userPrompt: String) async throws -> String {
|
|
|
|
|
+ try await completeText(systemPrompt: systemPrompt, userPrompt: userPrompt, maxTokens: nil)
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ func completeJSON<T: Decodable>(
|
|
|
|
|
+ systemPrompt: String,
|
|
|
|
|
+ userPrompt: String,
|
|
|
|
|
+ as type: T.Type
|
|
|
|
|
+ ) async throws -> T {
|
|
|
|
|
+ try await completeJSON(systemPrompt: systemPrompt, userPrompt: userPrompt, maxTokens: nil, as: type)
|
|
|
|
|
+ }
|
|
|
|
|
+}
|
|
|
|
|
+
|
|
|
struct OpenAIChatService: OpenAIChatServicing {
|
|
struct OpenAIChatService: OpenAIChatServicing {
|
|
|
static let shared = OpenAIChatService()
|
|
static let shared = OpenAIChatService()
|
|
|
|
|
|
|
@@ -43,16 +58,17 @@ struct OpenAIChatService: OpenAIChatServicing {
|
|
|
self.model = model
|
|
self.model = model
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
- func completeText(systemPrompt: String, userPrompt: String) async throws -> String {
|
|
|
|
|
- try await performRequest(systemPrompt: systemPrompt, userPrompt: userPrompt, jsonMode: false)
|
|
|
|
|
|
|
+ func completeText(systemPrompt: String, userPrompt: String, maxTokens: Int? = nil) async throws -> String {
|
|
|
|
|
+ try await performRequest(systemPrompt: systemPrompt, userPrompt: userPrompt, jsonMode: false, maxTokens: maxTokens)
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
func completeJSON<T: Decodable>(
|
|
func completeJSON<T: Decodable>(
|
|
|
systemPrompt: String,
|
|
systemPrompt: String,
|
|
|
userPrompt: String,
|
|
userPrompt: String,
|
|
|
|
|
+ maxTokens: Int? = nil,
|
|
|
as type: T.Type
|
|
as type: T.Type
|
|
|
) async throws -> T {
|
|
) async throws -> T {
|
|
|
- let content = try await performRequest(systemPrompt: systemPrompt, userPrompt: userPrompt, jsonMode: true)
|
|
|
|
|
|
|
+ let content = try await performRequest(systemPrompt: systemPrompt, userPrompt: userPrompt, jsonMode: true, maxTokens: maxTokens)
|
|
|
guard let data = content.data(using: .utf8) else {
|
|
guard let data = content.data(using: .utf8) else {
|
|
|
throw OpenAIServiceError.invalidResponse
|
|
throw OpenAIServiceError.invalidResponse
|
|
|
}
|
|
}
|
|
@@ -62,7 +78,8 @@ struct OpenAIChatService: OpenAIChatServicing {
|
|
|
private func performRequest(
|
|
private func performRequest(
|
|
|
systemPrompt: String,
|
|
systemPrompt: String,
|
|
|
userPrompt: String,
|
|
userPrompt: String,
|
|
|
- jsonMode: Bool
|
|
|
|
|
|
|
+ jsonMode: Bool,
|
|
|
|
|
+ maxTokens: Int?
|
|
|
) async throws -> String {
|
|
) async throws -> String {
|
|
|
guard let apiKey = apiKeyProvider() else {
|
|
guard let apiKey = apiKeyProvider() else {
|
|
|
throw OpenAIServiceError.missingAPIKey
|
|
throw OpenAIServiceError.missingAPIKey
|
|
@@ -80,7 +97,8 @@ struct OpenAIChatService: OpenAIChatServicing {
|
|
|
.init(role: "system", content: systemPrompt),
|
|
.init(role: "system", content: systemPrompt),
|
|
|
.init(role: "user", content: userPrompt)
|
|
.init(role: "user", content: userPrompt)
|
|
|
],
|
|
],
|
|
|
- responseFormat: jsonMode ? .init(type: "json_object") : nil
|
|
|
|
|
|
|
+ responseFormat: jsonMode ? .init(type: "json_object") : nil,
|
|
|
|
|
+ maxTokens: maxTokens
|
|
|
)
|
|
)
|
|
|
|
|
|
|
|
request.httpBody = try JSONEncoder().encode(body)
|
|
request.httpBody = try JSONEncoder().encode(body)
|
|
@@ -146,11 +164,13 @@ private struct ChatCompletionRequest: Encodable {
|
|
|
let model: String
|
|
let model: String
|
|
|
let messages: [Message]
|
|
let messages: [Message]
|
|
|
let responseFormat: ResponseFormat?
|
|
let responseFormat: ResponseFormat?
|
|
|
|
|
+ let maxTokens: Int?
|
|
|
|
|
|
|
|
enum CodingKeys: String, CodingKey {
|
|
enum CodingKeys: String, CodingKey {
|
|
|
case model
|
|
case model
|
|
|
case messages
|
|
case messages
|
|
|
case responseFormat = "response_format"
|
|
case responseFormat = "response_format"
|
|
|
|
|
+ case maxTokens = "max_tokens"
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
func encode(to encoder: Encoder) throws {
|
|
func encode(to encoder: Encoder) throws {
|
|
@@ -158,6 +178,7 @@ private struct ChatCompletionRequest: Encodable {
|
|
|
try container.encode(model, forKey: .model)
|
|
try container.encode(model, forKey: .model)
|
|
|
try container.encode(messages, forKey: .messages)
|
|
try container.encode(messages, forKey: .messages)
|
|
|
try container.encodeIfPresent(responseFormat, forKey: .responseFormat)
|
|
try container.encodeIfPresent(responseFormat, forKey: .responseFormat)
|
|
|
|
|
+ try container.encodeIfPresent(maxTokens, forKey: .maxTokens)
|
|
|
}
|
|
}
|
|
|
}
|
|
}
|
|
|
|
|
|