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.

473 lines
18 KiB

import Foundation
/// OpenAI服务异常
public class OpenAIError: Error {
let message: String
init(_ message: String) {
self.message = message
}
}
/// 工具调用信息
private class ToolCallInfo {
var id: String = ""
var name: String = ""
var arguments: String = ""
var isValid: Bool {
return !id.isEmpty && !name.isEmpty
}
}
/// OpenAI服务iOS原生实现
public class OpenAIService {
private let TAG = "OpenAIService"
private var baseUrl = "https://api.openai.com/v1/chat/completions"
private var apiKey: String = ""
private var isInitialized = false
private var model: String = "doubao-1-5-lite-32k-250115" // 默认模型
// 用于存储注册的函数
private var registeredFunctions: [[String: Any]] = []
// URL会话
private let session: URLSession
public init() {
// 创建URL会话配置
let config = URLSessionConfiguration.default
config.timeoutIntervalForRequest = 30.0
config.timeoutIntervalForResource = 30.0
session = URLSession(configuration: config)
}
/// 创建用户消息
public func createUserMessage(content: String) -> [String: Any] {
return ["role": "user", "content": content]
}
/// 创建助手消息
public func createAssistantMessage(content: String) -> [String: Any] {
return ["role": "assistant", "content": content]
}
/// 初始化OpenAI服务
public func initialize(apiKey: String, baseUrl: String = "", model: String = "") -> Bool {
self.apiKey = apiKey
if !baseUrl.isEmpty {
self.baseUrl = baseUrl
}
if !model.isEmpty {
self.model = model
}
isInitialized = !apiKey.isEmpty
return isInitialized
}
/// 注册函数
public func registerFunction(name: String, description: String, parameters: [String: Any]) -> Bool {
do {
let function: [String: Any] = [
"name": name,
"description": description,
"parameters": parameters
]
// 检查是否已存在相同名称的函数
if let existingIndex = registeredFunctions.firstIndex(where: { ($0["name"] as? String) == name }) {
// 如果已存在,则替换
registeredFunctions[existingIndex] = function
} else {
// 如果不存在,则添加
registeredFunctions.append(function)
}
return true
} catch {
return false
}
}
/// 发送消息(非流式输出)
public func sendMessage(messages: [[String: Any]], systemPrompt: String) throws -> String {
guard isInitialized, !apiKey.isEmpty else {
throw OpenAIError("OpenAI服务未初始化")
}
// 构建完整消息,添加系统提示
var fullMessages: [[String: Any]] = [
["role": "system", "content": systemPrompt]
]
fullMessages.append(contentsOf: messages)
// 构建请求体
var requestDict: [String: Any] = [
"model": model,
"messages": fullMessages,
"temperature": 0.7,
"max_tokens": 2000,
"stream": false
]
// 如果有注册的函数,添加到请求中
if !registeredFunctions.isEmpty {
var tools: [[String: Any]] = []
for function in registeredFunctions {
let tool: [String: Any] = [
"type": "function",
"function": function
]
tools.append(tool)
}
requestDict["tools"] = tools
}
// 将请求数据转换为JSON数据
guard let jsonData = try? JSONSerialization.data(withJSONObject: requestDict) else {
throw OpenAIError("无法序列化请求数据")
}
// 创建URL请求
guard let url = URL(string: baseUrl) else {
throw OpenAIError("无效的URL")
}
var request = URLRequest(url: url)
request.httpMethod = "POST"
request.addValue("application/json", forHTTPHeaderField: "Content-Type")
request.addValue("Bearer \(apiKey)", forHTTPHeaderField: "Authorization")
request.httpBody = jsonData
// 创建信号量用于同步请求
let semaphore = DispatchSemaphore(value: 0)
var responseResult: Result<String, Error> = .failure(OpenAIError("未收到响应"))
// 执行请求
let task = session.dataTask(with: request) { data, response, error in
if let error = error {
responseResult = .failure(OpenAIError("请求失败: \(error.localizedDescription)"))
semaphore.signal()
return
}
guard let httpResponse = response as? HTTPURLResponse else {
responseResult = .failure(OpenAIError("无效的HTTP响应"))
semaphore.signal()
return
}
guard httpResponse.statusCode == 200 else {
responseResult = .failure(OpenAIError("API调用失败: \(httpResponse.statusCode)"))
semaphore.signal()
return
}
guard let data = data else {
responseResult = .failure(OpenAIError("响应数据为空"))
semaphore.signal()
return
}
do {
// 解析JSON响应
guard let jsonResponse = try JSONSerialization.jsonObject(with: data) as? [String: Any] else {
responseResult = .failure(OpenAIError("无法解析JSON响应"))
semaphore.signal()
return
}
// 检查是否有函数调用
if let choices = jsonResponse["choices"] as? [[String: Any]], !choices.isEmpty,
let choice = choices.first,
let message = choice["message"] as? [String: Any] {
// 检查是否有工具调用
if let toolCalls = message["tool_calls"] as? [[String: Any]], !toolCalls.isEmpty,
let toolCall = toolCalls.first,
let function = toolCall["function"] as? [String: Any],
let name = function["name"] as? String,
let arguments = function["arguments"] as? String,
let id = toolCall["id"] as? String {
let functionCallDict: [String: Any] = [
"name": name,
"arguments": arguments,
"id": id
]
// 将函数调用转为JSON字符串
if let functionCallData = try? JSONSerialization.data(withJSONObject: functionCallDict),
let functionCallString = String(data: functionCallData, encoding: .utf8) {
responseResult = .success(functionCallString)
semaphore.signal()
return
}
}
// 如果没有工具调用,返回消息内容
if let content = message["content"] as? String {
responseResult = .success(content)
semaphore.signal()
return
}
}
responseResult = .failure(OpenAIError("无效的响应格式"))
semaphore.signal()
} catch {
responseResult = .failure(OpenAIError("解析响应时出错: \(error.localizedDescription)"))
semaphore.signal()
}
}
task.resume()
// 等待响应完成
_ = semaphore.wait(timeout: .distantFuture)
// 返回结果或抛出错误
switch responseResult {
case .success(let result):
return result
case .failure(let error):
throw error
}
}
/// 发送消息(流式输出)
public func sendMessageStream(messages: [[String: Any]], systemPrompt: String, callback: @escaping StreamCallback) {
guard isInitialized, !apiKey.isEmpty else {
callback.onError(OpenAIError("OpenAI服务未初始化"))
return
}
// 构建完整消息,添加系统提示
var fullMessages: [[String: Any]] = [
["role": "system", "content": systemPrompt]
]
fullMessages.append(contentsOf: messages)
// 构建请求体
var requestDict: [String: Any] = [
"model": model,
"messages": fullMessages,
"temperature": 0.7,
"max_tokens": 2000,
"stream": true
]
// 如果有注册的函数,添加到请求中
if !registeredFunctions.isEmpty {
var tools: [[String: Any]] = []
for function in registeredFunctions {
let tool: [String: Any] = [
"type": "function",
"function": function
]
tools.append(tool)
}
requestDict["tools"] = tools
}
// 将请求数据转换为JSON数据
guard let jsonData = try? JSONSerialization.data(withJSONObject: requestDict) else {
callback.onError(OpenAIError("无法序列化请求数据"))
return
}
// 创建URL请求
guard let url = URL(string: baseUrl) else {
callback.onError(OpenAIError("无效的URL"))
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")
request.httpBody = jsonData
// 用于存储函数调用的各个部分
var finalToolCalls: [Int: ToolCallInfo] = [:]
// 创建数据任务
let task = session.dataTask(with: request) { data, response, error in
if let error = error {
callback.onError(OpenAIError("请求失败: \(error.localizedDescription)"))
return
}
guard let httpResponse = response as? HTTPURLResponse else {
callback.onError(OpenAIError("无效的HTTP响应"))
return
}
guard httpResponse.statusCode == 200 else {
callback.onError(OpenAIError("API调用失败: \(httpResponse.statusCode)"))
return
}
guard let data = data else {
callback.onError(OpenAIError("响应数据为空"))
return
}
// 处理SSE数据流
if let text = String(data: data, encoding: .utf8) {
let lines = text.components(separatedBy: "\n")
for line in lines {
if line.isEmpty { continue }
if line.hasPrefix("data: ") {
let dataContent = line.dropFirst(6)
// 处理[DONE]消息
if dataContent == "[DONE]" {
self.processToolCalls(finalToolCalls, callback: callback)
callback.onComplete()
break
}
// 解析JSON数据
do {
if let data = dataContent.data(using: .utf8),
let jsonData = try JSONSerialization.jsonObject(with: data) as? [String: Any] {
// 处理消息内容
if let choices = jsonData["choices"] as? [[String: Any]], !choices.isEmpty,
let choice = choices.first {
if let delta = choice["delta"] as? [String: Any] {
// 处理普通文本内容
if let content = delta["content"] as? String {
callback.onToken(content)
}
// 处理工具调用(函数调用)
if let toolCalls = delta["tool_calls"] as? [[String: Any]] {
for toolCall in toolCalls {
if let index = toolCall["index"] as? Int {
// 创建或获取现有的工具调用信息
let toolCallInfo = finalToolCalls[index] ?? ToolCallInfo()
// 更新ID
if let id = toolCall["id"] as? String {
toolCallInfo.id = id
}
// 更新函数信息
if let function = toolCall["function"] as? [String: Any] {
if let name = function["name"] as? String {
toolCallInfo.name = name
}
if let arguments = function["arguments"] as? String {
toolCallInfo.arguments += arguments
}
}
finalToolCalls[index] = toolCallInfo
}
}
}
}
}
}
} catch {
NSLog("解析JSON出错: \(error.localizedDescription)")
// 忽略解析错误,继续处理其他行
}
}
}
}
}
task.resume()
}
/// 发送函数调用结果
public func sendFunctionCallResult(
messages: [[String: Any]],
systemPrompt: String,
functionCall: [String: Any],
functionResult: String,
callback: @escaping StreamCallback
) {
do {
// 构建完整消息数组
var fullMessages: [[String: Any]] = [
// 添加系统提示
["role": "system", "content": systemPrompt]
]
// 添加用户消息
fullMessages.append(contentsOf: messages)
// 获取函数相关信息
guard let name = functionCall["name"] as? String,
let arguments = functionCall["arguments"] as? String else {
callback.onError(OpenAIError("函数调用信息不完整"))
return
}
let id = functionCall["id"] as? String ?? "call_\(Int(Date().timeIntervalSince1970 * 1000))"
// 添加函数调用消息
fullMessages.append([
"role": "assistant",
"content": NSNull(),
"tool_calls": [
[
"id": id,
"type": "function",
"function": [
"name": name,
"arguments": arguments
]
]
]
])
// 添加函数调用结果
fullMessages.append([
"role": "tool",
"content": functionResult,
"tool_call_id": id
])
// 发送完整对话
sendMessageStream(messages: fullMessages, systemPrompt: systemPrompt, callback: callback)
} catch {
callback.onError(OpenAIError("发送函数调用结果失败: \(error.localizedDescription)"))
}
}
/// 处理工具调用(函数调用)并回调
private func processToolCalls(_ toolCalls: [Int: ToolCallInfo], callback: StreamCallback) {
if toolCalls.isEmpty { return }
// 只处理第一个工具调用
guard let firstToolCall = toolCalls.values.first, firstToolCall.isValid else { return }
// 创建函数调用字典
let functionCall: [String: Any] = [
"name": firstToolCall.name,
"arguments": firstToolCall.arguments,
"id": firstToolCall.id
]
// 回调
callback.onFunctionCall(functionCall)
}
/// 流式输出回调协议
public typealias StreamCallback = (onToken: (String) -> Void,
onComplete: () -> Void,
onError: (Error) -> Void,
onFunctionCall: ([String: Any]) -> Void)
}