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

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)"
}
}
}