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.
531 lines
20 KiB
531 lines
20 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 requestBody: [String: Any] = [
|
|
"model": model,
|
|
"temperature": 0.7,
|
|
"max_tokens": 2000,
|
|
"stream": true
|
|
]
|
|
|
|
// 构建完整消息数组,添加系统提示
|
|
var fullMessages: [[String: Any]] = [
|
|
["role": "system", "content": systemPrompt]
|
|
]
|
|
fullMessages.append(contentsOf: messages)
|
|
requestBody["messages"] = fullMessages
|
|
|
|
// 添加工具列表
|
|
if let toolMaps = mcpClient?.getToolMaps(), !toolMaps.isEmpty {
|
|
var tools: [[String: Any]] = []
|
|
|
|
for toolMap in toolMaps {
|
|
if let tool = toolMap as? [String: Any] {
|
|
tools.append(tool)
|
|
}
|
|
}
|
|
|
|
if !tools.isEmpty {
|
|
requestBody["tools"] = tools
|
|
}
|
|
} else if !registeredFunctions.isEmpty {
|
|
var tools: [[String: Any]] = []
|
|
|
|
for function in registeredFunctions {
|
|
let tool: [String: Any] = [
|
|
"type": "function",
|
|
"function": function
|
|
]
|
|
tools.append(tool)
|
|
}
|
|
|
|
if !tools.isEmpty {
|
|
requestBody["tools"] = tools
|
|
}
|
|
}
|
|
|
|
// 转换为JSON数据
|
|
guard let jsonData = try? JSONSerialization.data(withJSONObject: requestBody) 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
|
|
|
|
// 创建流式会话任务
|
|
let delegate = SSEStreamDelegate(callback: callback)
|
|
let session = URLSession(configuration: .default, delegate: delegate, delegateQueue: nil)
|
|
let task = session.dataTask(with: request)
|
|
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)"))
|
|
}
|
|
}
|
|
|
|
/// SSE流委托实现
|
|
private class SSEStreamDelegate: NSObject, URLSessionDataDelegate {
|
|
let callback: StreamCallback
|
|
private var buffer = Data()
|
|
private var finalToolCalls: [Int: ToolCallInfo] = [:]
|
|
|
|
init(callback: @escaping StreamCallback) {
|
|
self.callback = callback
|
|
super.init()
|
|
}
|
|
|
|
// 接收数据流
|
|
func urlSession(_ session: URLSession, dataTask: URLSessionDataTask, didReceive data: Data) {
|
|
buffer.append(data)
|
|
|
|
// 处理可能包含多行的数据
|
|
processBuffer()
|
|
}
|
|
|
|
// 处理缓冲区数据
|
|
private func processBuffer() {
|
|
// 按行分割
|
|
while let newlineIndex = buffer.firstIndex(of: 10) { // 10是换行符的ASCII码
|
|
let lineData = buffer.prefix(upTo: newlineIndex)
|
|
buffer.removeSubrange(0...newlineIndex) // 移除已处理的行,包括换行符
|
|
|
|
// 解析行数据
|
|
if let line = String(data: lineData, encoding: .utf8)?.trimmingCharacters(in: .whitespacesAndNewlines) {
|
|
processLine(line)
|
|
}
|
|
}
|
|
}
|
|
|
|
// 处理单行数据
|
|
private func processLine(_ line: String) {
|
|
guard !line.isEmpty else { return }
|
|
|
|
if line.hasPrefix("data: ") {
|
|
let dataContent = line.dropFirst(6)
|
|
|
|
// 处理[DONE]消息
|
|
if dataContent == "[DONE]" {
|
|
let hasToolCalls = self.processToolCalls()
|
|
// 只有在没有工具调用时才认为对话真正完成
|
|
if !hasToolCalls {
|
|
callback.onComplete()
|
|
}
|
|
return
|
|
}
|
|
|
|
// 解析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)")
|
|
// 忽略解析错误,继续处理其他行
|
|
}
|
|
}
|
|
}
|
|
|
|
// 处理工具调用
|
|
private func processToolCalls() -> Bool {
|
|
if finalToolCalls.isEmpty { return false }
|
|
|
|
// 只处理第一个工具调用
|
|
guard let firstToolCall = finalToolCalls.values.first, firstToolCall.isValid else { return false }
|
|
|
|
// 创建函数调用字典
|
|
let functionCall: [String: Any] = [
|
|
"name": firstToolCall.name,
|
|
"arguments": firstToolCall.arguments,
|
|
"id": firstToolCall.id
|
|
]
|
|
|
|
// 回调
|
|
callback.onFunctionCall(functionCall)
|
|
return true
|
|
}
|
|
|
|
// 处理完成
|
|
func urlSession(_ session: URLSession, task: URLSessionTask, didCompleteWithError error: Error?) {
|
|
if let error = error {
|
|
callback.onError(OpenAIError("请求失败: \(error.localizedDescription)"))
|
|
}
|
|
}
|
|
}
|
|
|
|
/// 流式输出回调协议
|
|
public typealias StreamCallback = (onToken: (String) -> Void,
|
|
onComplete: () -> Void,
|
|
onError: (Error) -> Void,
|
|
onFunctionCall: ([String: Any]) -> Void)
|
|
|
|
/// 处理工具调用(函数调用)并回调
|
|
private func processToolCalls(_ toolCalls: [Int: ToolCallInfo], callback: StreamCallback) -> Bool {
|
|
if toolCalls.isEmpty { return false }
|
|
|
|
// 只处理第一个工具调用
|
|
guard let firstToolCall = toolCalls.values.first, firstToolCall.isValid else { return false }
|
|
|
|
// 创建函数调用字典
|
|
let functionCall: [String: Any] = [
|
|
"name": firstToolCall.name,
|
|
"arguments": firstToolCall.arguments,
|
|
"id": firstToolCall.id
|
|
]
|
|
|
|
// 回调
|
|
callback.onFunctionCall(functionCall)
|
|
return true
|
|
}
|
|
}
|