You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
391 lines
13 KiB
391 lines
13 KiB
import Foundation
|
|
|
|
/**
|
|
* 火山AI服务的iOS原生实现
|
|
*
|
|
* 参考Android端的VolcanoAIService实现,提供同步和异步的API调用方式
|
|
*/
|
|
class VolcanoAIService {
|
|
private let TAG = "VolcanoAIService"
|
|
private let baseUrl = "https://ark.cn-beijing.volces.com/api/v3"
|
|
private let chatEndpoint = "/chat/completions"
|
|
private let session: URLSession
|
|
|
|
private var apiKey: String = ""
|
|
private var isInitialized = false
|
|
|
|
init() {
|
|
let config = URLSessionConfiguration.default
|
|
config.timeoutIntervalForRequest = 30.0
|
|
config.timeoutIntervalForResource = 30.0
|
|
self.session = URLSession(configuration: config)
|
|
}
|
|
|
|
/**
|
|
* 初始化火山AI服务
|
|
*
|
|
* @param apiKey 火山AI API密钥
|
|
* @return 初始化是否成功
|
|
*/
|
|
func initialize(apiKey: String) -> Bool {
|
|
self.apiKey = apiKey
|
|
isInitialized = !apiKey.isEmpty
|
|
|
|
if !isInitialized {
|
|
NSLog("%@: 初始化失败:API key 不能为空", TAG)
|
|
} else {
|
|
NSLog("%@: 火山AI服务初始化成功", TAG)
|
|
}
|
|
|
|
return isInitialized
|
|
}
|
|
|
|
/**
|
|
* 生成个性化问候语
|
|
*
|
|
* @param agentName 代理名称
|
|
* @param systemPrompt 系统提示词
|
|
* @param callback 回调函数,返回生成的问候语
|
|
*/
|
|
func generateGreeting(agentName: String, systemPrompt: String, callback: @escaping (String?, Error?) -> Void) {
|
|
let messages: [[String: Any]] = [
|
|
["role": "system", "content": systemPrompt],
|
|
["role": "user", "content": "请用一句简短的话向我打个招呼,要符合你的身份和性格特点,不要超过18个字。"]
|
|
]
|
|
|
|
let messagesData = try? JSONSerialization.data(withJSONObject: messages, options: [])
|
|
let messagesArray = try? JSONSerialization.jsonObject(with: messagesData!, options: []) as? [[String: Any]]
|
|
|
|
var result = ""
|
|
|
|
sendMessageStream(messages: messagesArray!, systemPrompt: systemPrompt, streamCallback: StreamCallback(
|
|
onToken: { token in
|
|
result.append(token)
|
|
},
|
|
onComplete: {
|
|
callback(result, nil)
|
|
},
|
|
onError: { error in
|
|
callback(nil, error)
|
|
}
|
|
))
|
|
}
|
|
|
|
/**
|
|
* 发送消息(非流式输出)
|
|
*
|
|
* @param messages 消息列表
|
|
* @param systemPrompt 系统提示词
|
|
* @return 返回AI的回复
|
|
* @throws VolcanoAIError 如果API调用失败
|
|
*/
|
|
func sendMessage(messages: [[String: Any]], systemPrompt: String) throws -> String {
|
|
// 检查是否已初始化
|
|
if !isInitialized || apiKey.isEmpty {
|
|
throw VolcanoAIError.serviceNotInitialized
|
|
}
|
|
|
|
var fullMessages: [[String: Any]] = [
|
|
["role": "system", "content": systemPrompt]
|
|
]
|
|
|
|
fullMessages.append(contentsOf: messages)
|
|
|
|
let requestBody: [String: Any] = [
|
|
"model": "doubao-1-5-lite-32k-250115",
|
|
"messages": fullMessages,
|
|
"temperature": 0.7,
|
|
"max_tokens": 2000,
|
|
"stream": false
|
|
]
|
|
|
|
guard let url = URL(string: "\(baseUrl)\(chatEndpoint)") else {
|
|
throw VolcanoAIError.invalidURL
|
|
}
|
|
|
|
var request = URLRequest(url: url)
|
|
request.httpMethod = "POST"
|
|
request.addValue("application/json", forHTTPHeaderField: "Content-Type")
|
|
request.addValue("Bearer \(apiKey)", forHTTPHeaderField: "Authorization")
|
|
|
|
do {
|
|
request.httpBody = try JSONSerialization.data(withJSONObject: requestBody, options: [])
|
|
} catch {
|
|
throw VolcanoAIError.invalidRequestBody
|
|
}
|
|
|
|
let semaphore = DispatchSemaphore(value: 0)
|
|
var responseData: Data?
|
|
var responseError: Error?
|
|
|
|
let task = session.dataTask(with: request) { data, response, error in
|
|
if let error = error {
|
|
responseError = VolcanoAIError.networkError(error.localizedDescription)
|
|
semaphore.signal()
|
|
return
|
|
}
|
|
|
|
guard let httpResponse = response as? HTTPURLResponse else {
|
|
responseError = VolcanoAIError.invalidResponse
|
|
semaphore.signal()
|
|
return
|
|
}
|
|
|
|
if !(200...299).contains(httpResponse.statusCode) {
|
|
var errorMessage = "Unknown error occurred"
|
|
if let data = data, let json = try? JSONSerialization.jsonObject(with: data) as? [String: Any],
|
|
let error = json["error"] as? [String: Any],
|
|
let message = error["message"] as? String {
|
|
errorMessage = message
|
|
}
|
|
responseError = VolcanoAIError.apiError(errorMessage)
|
|
semaphore.signal()
|
|
return
|
|
}
|
|
|
|
responseData = data
|
|
semaphore.signal()
|
|
}
|
|
|
|
task.resume()
|
|
_ = semaphore.wait(timeout: .distantFuture)
|
|
|
|
if let error = responseError {
|
|
throw error
|
|
}
|
|
|
|
guard let data = responseData else {
|
|
throw VolcanoAIError.emptyResponse
|
|
}
|
|
|
|
do {
|
|
guard let json = try JSONSerialization.jsonObject(with: data) as? [String: Any],
|
|
let choices = json["choices"] as? [[String: Any]],
|
|
let firstChoice = choices.first,
|
|
let message = firstChoice["message"] as? [String: Any],
|
|
let content = message["content"] as? String else {
|
|
throw VolcanoAIError.invalidResponseFormat
|
|
}
|
|
|
|
return content
|
|
} catch {
|
|
throw VolcanoAIError.invalidResponseFormat
|
|
}
|
|
}
|
|
|
|
/**
|
|
* 发送消息(流式输出)
|
|
*
|
|
* @param messages 消息列表
|
|
* @param systemPrompt 系统提示词
|
|
* @param streamCallback 回调函数,用于接收流式输出的结果
|
|
*/
|
|
func sendMessageStream(messages: [[String: Any]], systemPrompt: String, streamCallback: StreamCallback) {
|
|
// 检查是否已初始化
|
|
if !isInitialized || apiKey.isEmpty {
|
|
streamCallback.onError(VolcanoAIError.serviceNotInitialized)
|
|
return
|
|
}
|
|
|
|
var fullMessages: [[String: Any]] = [
|
|
["role": "system", "content": systemPrompt]
|
|
]
|
|
|
|
fullMessages.append(contentsOf: messages)
|
|
|
|
let requestBody: [String: Any] = [
|
|
"model": "doubao-1-5-lite-32k-250115",
|
|
"messages": fullMessages,
|
|
"temperature": 0.7,
|
|
"max_tokens": 2000,
|
|
"stream": true
|
|
]
|
|
|
|
guard let url = URL(string: "\(baseUrl)\(chatEndpoint)") else {
|
|
streamCallback.onError(VolcanoAIError.invalidURL)
|
|
return
|
|
}
|
|
|
|
var request = URLRequest(url: url)
|
|
request.httpMethod = "POST"
|
|
request.addValue("application/json", forHTTPHeaderField: "Content-Type")
|
|
request.addValue("Bearer \(apiKey)", forHTTPHeaderField: "Authorization")
|
|
request.addValue("text/event-stream", forHTTPHeaderField: "Accept")
|
|
|
|
do {
|
|
request.httpBody = try JSONSerialization.data(withJSONObject: requestBody, options: [])
|
|
} catch {
|
|
streamCallback.onError(VolcanoAIError.invalidRequestBody)
|
|
return
|
|
}
|
|
|
|
let task = session.dataTask(with: request) { data, response, error in
|
|
if let error = error {
|
|
streamCallback.onError(VolcanoAIError.networkError(error.localizedDescription))
|
|
return
|
|
}
|
|
|
|
guard let httpResponse = response as? HTTPURLResponse else {
|
|
streamCallback.onError(VolcanoAIError.invalidResponse)
|
|
return
|
|
}
|
|
|
|
if !(200...299).contains(httpResponse.statusCode) {
|
|
var errorMessage = "Unknown error occurred"
|
|
if let data = data, let json = try? JSONSerialization.jsonObject(with: data) as? [String: Any],
|
|
let error = json["error"] as? [String: Any],
|
|
let message = error["message"] as? String {
|
|
errorMessage = message
|
|
}
|
|
streamCallback.onError(VolcanoAIError.apiError(errorMessage))
|
|
return
|
|
}
|
|
|
|
guard let data = data else {
|
|
streamCallback.onError(VolcanoAIError.emptyResponse)
|
|
return
|
|
}
|
|
|
|
// 处理SSE流数据
|
|
let responseString = String(data: data, encoding: .utf8) ?? ""
|
|
let lines = responseString.components(separatedBy: "\n")
|
|
|
|
for line in lines {
|
|
if line.isEmpty { continue }
|
|
|
|
if line.hasPrefix("data: ") {
|
|
let data = String(line.dropFirst(6))
|
|
if data == "[DONE]" {
|
|
streamCallback.onComplete()
|
|
break
|
|
}
|
|
|
|
do {
|
|
if let jsonData = data.data(using: .utf8),
|
|
let json = try JSONSerialization.jsonObject(with: jsonData) as? [String: Any],
|
|
let choices = json["choices"] as? [[String: Any]],
|
|
let firstChoice = choices.first,
|
|
let delta = firstChoice["delta"] as? [String: Any],
|
|
let content = delta["content"] as? String {
|
|
streamCallback.onToken(content)
|
|
}
|
|
} catch {
|
|
// 忽略无效的JSON数据
|
|
continue
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
task.resume()
|
|
}
|
|
|
|
/**
|
|
* 同步方式发送消息(流式输出)
|
|
*
|
|
* 注意:此方法会阻塞当前线程,请在后台线程中调用
|
|
*
|
|
* @param messages 消息列表
|
|
* @param systemPrompt 系统提示词
|
|
* @return 返回完整的AI回复
|
|
* @throws VolcanoAIError 如果API调用失败
|
|
*/
|
|
func sendMessageStreamSync(messages: [[String: Any]], systemPrompt: String) throws -> String {
|
|
var result = ""
|
|
let semaphore = DispatchSemaphore(value: 0)
|
|
var responseError: Error?
|
|
|
|
sendMessageStream(messages: messages, systemPrompt: systemPrompt, streamCallback: StreamCallback(
|
|
onToken: { token in
|
|
result.append(token)
|
|
},
|
|
onComplete: {
|
|
semaphore.signal()
|
|
},
|
|
onError: { error in
|
|
responseError = error
|
|
semaphore.signal()
|
|
}
|
|
))
|
|
|
|
// 等待流式输出完成或出错
|
|
_ = semaphore.wait(timeout: .now() + 60)
|
|
|
|
if let error = responseError {
|
|
throw error
|
|
}
|
|
|
|
return result
|
|
}
|
|
|
|
/**
|
|
* 创建用户消息
|
|
*/
|
|
func createUserMessage(content: String) -> [String: Any] {
|
|
return ["role": "user", "content": content]
|
|
}
|
|
|
|
/**
|
|
* 创建系统消息
|
|
*/
|
|
func createSystemMessage(content: String) -> [String: Any] {
|
|
return ["role": "system", "content": content]
|
|
}
|
|
|
|
/**
|
|
* 创建助手消息
|
|
*/
|
|
func createAssistantMessage(content: String) -> [String: Any] {
|
|
return ["role": "assistant", "content": content]
|
|
}
|
|
|
|
/**
|
|
* 流式输出回调类
|
|
*/
|
|
class StreamCallback {
|
|
let onToken: (String) -> Void
|
|
let onComplete: () -> Void
|
|
let onError: (Error) -> Void
|
|
|
|
init(onToken: @escaping (String) -> Void, onComplete: @escaping () -> Void, onError: @escaping (Error) -> Void) {
|
|
self.onToken = onToken
|
|
self.onComplete = onComplete
|
|
self.onError = onError
|
|
}
|
|
}
|
|
}
|
|
|
|
/**
|
|
* 火山AI错误枚举
|
|
*/
|
|
enum VolcanoAIError: Error {
|
|
case serviceNotInitialized
|
|
case invalidURL
|
|
case invalidRequestBody
|
|
case networkError(String)
|
|
case invalidResponse
|
|
case emptyResponse
|
|
case invalidResponseFormat
|
|
case apiError(String)
|
|
|
|
var localizedDescription: String {
|
|
switch self {
|
|
case .serviceNotInitialized:
|
|
return "火山AI服务未初始化或API key为空,请先调用initialize方法"
|
|
case .invalidURL:
|
|
return "无效的URL"
|
|
case .invalidRequestBody:
|
|
return "无效的请求体"
|
|
case .networkError(let message):
|
|
return "网络错误: \(message)"
|
|
case .invalidResponse:
|
|
return "无效的响应"
|
|
case .emptyResponse:
|
|
return "空响应"
|
|
case .invalidResponseFormat:
|
|
return "无效的响应格式"
|
|
case .apiError(let message):
|
|
return "API错误: \(message)"
|
|
}
|
|
}
|
|
}
|