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.
830 lines
31 KiB
830 lines
31 KiB
import Foundation
|
|
import UIKit
|
|
|
|
/// 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服务流回调结构体定义
|
|
struct StreamCallback {
|
|
let onToken: (String) -> Void
|
|
let onComplete: () -> Void
|
|
let onError: (Error) -> Void
|
|
let onFunctionCall: ([String: Any]) -> Void
|
|
let onFunctionCallResult: ([String: Any], [String: Any]) -> Void
|
|
|
|
init(
|
|
onToken: @escaping (String) -> Void,
|
|
onComplete: @escaping () -> Void,
|
|
onError: @escaping (Error) -> Void,
|
|
onFunctionCall: @escaping ([String: Any]) -> Void,
|
|
onFunctionCallResult: @escaping ([String: Any], [String: Any]) -> Void
|
|
) {
|
|
self.onToken = onToken
|
|
self.onComplete = onComplete
|
|
self.onError = onError
|
|
self.onFunctionCall = onFunctionCall
|
|
self.onFunctionCallResult = onFunctionCallResult
|
|
}
|
|
}
|
|
|
|
/// 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 visionModel: String = "doubao-1-5-vision-pro-32k-250115" // 默认视觉模型
|
|
|
|
// 用于存储注册的函数
|
|
private var registeredFunctions: [[String: Any]] = []
|
|
|
|
// URL会话
|
|
private let session: URLSession
|
|
|
|
// 当前流式请求任务
|
|
private var currentStreamTask: URLSessionDataTask?
|
|
private var isCanceled = false
|
|
|
|
// MCP客户端
|
|
private var mcpClient: MCPClient?
|
|
private var isMcpInitialized = false
|
|
|
|
// 是否自动处理MCP工具调用
|
|
private var autoHandleMcpTools = true
|
|
|
|
public init() {
|
|
// 创建URL会话配置
|
|
let config = URLSessionConfiguration.default
|
|
config.timeoutIntervalForRequest = 30.0
|
|
config.timeoutIntervalForResource = 30.0
|
|
session = URLSession(configuration: config)
|
|
}
|
|
|
|
/// 初始化MCP客户端
|
|
/// - Parameter serverUrl: MCP服务器地址
|
|
/// - Returns: 是否成功开始初始化
|
|
public func initializeMcpClient(_ serverUrl: String) -> Bool {
|
|
if mcpClient != nil {
|
|
mcpClient?.close()
|
|
}
|
|
|
|
mcpClient = MCPClient()
|
|
|
|
// 开始异步连接
|
|
DispatchQueue.global(qos: .userInitiated).async { [weak self] in
|
|
guard let self = self, let mcpClient = self.mcpClient else { return }
|
|
|
|
let result = mcpClient.connectToSSE(serverUrl)
|
|
self.isMcpInitialized = result
|
|
print("\(self.TAG) MCP客户端初始化\(result ? "成功" : "失败")")
|
|
}
|
|
|
|
return true // 立即返回,实际连接在后台进行
|
|
}
|
|
|
|
/// 检查MCP客户端是否已初始化
|
|
/// - Returns: 是否已初始化
|
|
public func checkMcpInitialized() -> Bool {
|
|
return isMcpInitialized && mcpClient?.checkIsConnected() == true
|
|
}
|
|
|
|
/// 关闭MCP客户端
|
|
/// - Returns: 是否成功关闭
|
|
public func closeMcpClient() -> Bool {
|
|
mcpClient?.close()
|
|
mcpClient = nil
|
|
isMcpInitialized = false
|
|
return true
|
|
}
|
|
|
|
/// 处理MCP工具调用
|
|
/// - Parameter functionCall: 函数调用信息
|
|
/// - Returns: 处理结果
|
|
public func handleMcpToolCall(_ functionCall: [String: Any]) async -> String {
|
|
if mcpClient == nil || !isMcpInitialized {
|
|
return "MCP客户端未初始化"
|
|
}
|
|
|
|
do {
|
|
// 获取函数名称
|
|
guard let name = functionCall["name"] as? String,
|
|
let argumentsJson = functionCall["arguments"] as? String else {
|
|
return "函数调用信息不完整"
|
|
}
|
|
|
|
// 解析参数
|
|
let arguments = mcpClient?.parseJsonArguments(argumentsJson) ?? [:]
|
|
|
|
// 调用工具
|
|
let result = await mcpClient?.callTool(name: name, arguments: arguments)
|
|
return result?["context"] as? String ?? "处理MCP工具调用失败"
|
|
} catch {
|
|
print("\(TAG) 处理MCP工具调用失败: \(error.localizedDescription)")
|
|
return "处理MCP工具调用失败: \(error.localizedDescription)"
|
|
}
|
|
}
|
|
|
|
/// 创建用户消息
|
|
public func createUserMessage(content: String) -> [String: Any] {
|
|
return ["role": "user", "content": content]
|
|
}
|
|
|
|
/// 创建系统消息
|
|
public func createSystemMessage(content: String) -> [String: Any] {
|
|
return ["role": "system", "content": content]
|
|
}
|
|
|
|
/// 创建助手消息
|
|
public func createAssistantMessage(content: String) -> [String: Any] {
|
|
return ["role": "assistant", "content": content]
|
|
}
|
|
|
|
/// 创建带图片的用户消息
|
|
public func createUserMessageWithImage(text: String, imageBase64: String) -> [String: Any] {
|
|
// 创建包含文本和图片的内容数组
|
|
var contentArray: [[String: Any]] = []
|
|
|
|
// 添加文本部分
|
|
if !text.isEmpty {
|
|
contentArray.append([
|
|
"type": "text",
|
|
"text": text
|
|
])
|
|
}
|
|
|
|
// 添加图片部分
|
|
contentArray.append([
|
|
"type": "image_url",
|
|
"image_url": [
|
|
"url": "data:image/jpeg;base64,\(imageBase64)"
|
|
]
|
|
])
|
|
|
|
return ["role": "user", "content": contentArray]
|
|
}
|
|
|
|
/// 将文件转换为Base64字符串
|
|
public func fileToBase64(_ filePath: String, maxSizeKB: Int = 20480) -> String? {
|
|
do {
|
|
let fileURL = URL(fileURLWithPath: filePath)
|
|
|
|
// 检查文件是否存在
|
|
guard FileManager.default.fileExists(atPath: filePath) else {
|
|
print("\(TAG) 文件不存在: \(filePath)")
|
|
return nil
|
|
}
|
|
|
|
// 读取图片
|
|
guard let originalImage = UIImage(contentsOfFile: filePath) else {
|
|
print("\(TAG) 无法读取图片: \(filePath)")
|
|
return nil
|
|
}
|
|
|
|
// 检查图片尺寸,限制最大为1024*1024
|
|
var processedImage = originalImage
|
|
let maxDimension: CGFloat = 1024
|
|
if originalImage.size.width > maxDimension || originalImage.size.height > maxDimension {
|
|
print("\(TAG) 图片尺寸超过限制,进行缩放: \(originalImage.size) -> \(maxDimension)")
|
|
|
|
// 计算缩放比例,保持纵横比
|
|
let widthRatio = maxDimension / originalImage.size.width
|
|
let heightRatio = maxDimension / originalImage.size.height
|
|
let ratio = min(widthRatio, heightRatio)
|
|
|
|
let newWidth = originalImage.size.width * ratio
|
|
let newHeight = originalImage.size.height * ratio
|
|
|
|
let newSize = CGSize(width: newWidth, height: newHeight)
|
|
UIGraphicsBeginImageContextWithOptions(newSize, false, 1.0)
|
|
originalImage.draw(in: CGRect(origin: .zero, size: newSize))
|
|
if let resizedImage = UIGraphicsGetImageFromCurrentImageContext() {
|
|
processedImage = resizedImage
|
|
}
|
|
UIGraphicsEndImageContext()
|
|
|
|
print("\(TAG) 缩放后图片尺寸: \(newSize)")
|
|
}
|
|
|
|
// 压缩图片
|
|
var imageData = processedImage.jpegData(compressionQuality: 0.9)
|
|
|
|
// 检查文件大小,如果超出限制,继续压缩
|
|
var compressionQuality: CGFloat = 0.9
|
|
while let data = imageData, data.count > maxSizeKB * 1024 && compressionQuality > 0.1 {
|
|
compressionQuality -= 0.1
|
|
imageData = processedImage.jpegData(compressionQuality: compressionQuality)
|
|
}
|
|
|
|
guard let finalImageData = imageData else {
|
|
print("\(TAG) 无法压缩图片")
|
|
return nil
|
|
}
|
|
|
|
// 检查最终大小
|
|
if finalImageData.count > maxSizeKB * 1024 {
|
|
print("\(TAG) 压缩后图片仍然超出大小限制: \(finalImageData.count / 1024)KB > \(maxSizeKB)KB")
|
|
}
|
|
|
|
// 转为Base64
|
|
return finalImageData.base64EncodedString()
|
|
} catch {
|
|
print("\(TAG) 转换文件到Base64失败: \(error.localizedDescription)")
|
|
return nil
|
|
}
|
|
}
|
|
|
|
/// 初始化OpenAI服务
|
|
public func initialize(apiKey: String, baseUrl: String = "", model: String = "", mcpServer: String = "") -> Bool {
|
|
self.apiKey = apiKey
|
|
if !baseUrl.isEmpty {
|
|
self.baseUrl = baseUrl
|
|
}
|
|
if !model.isEmpty {
|
|
self.model = model
|
|
}
|
|
|
|
// 初始化MCP客户端
|
|
if !mcpServer.isEmpty {
|
|
initializeMcpClient(mcpServer)
|
|
}
|
|
|
|
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 cancelCurrentStream() {
|
|
isCanceled = true
|
|
currentStreamTask?.cancel()
|
|
currentStreamTask = nil
|
|
}
|
|
|
|
/// 自动处理MCP工具调用
|
|
private func autoHandleMcpToolCall(
|
|
functionCall: [String: Any],
|
|
messages: [[String: Any]],
|
|
systemPrompt: String,
|
|
callback: StreamCallback
|
|
) {
|
|
Task {
|
|
do {
|
|
// 获取函数名称
|
|
guard let name = functionCall["name"] as? String,
|
|
let argumentsJson = functionCall["arguments"] as? String else {
|
|
callback.onError(OpenAIError("函数调用信息不完整"))
|
|
return
|
|
}
|
|
|
|
// 解析参数
|
|
let arguments = mcpClient?.parseJsonArguments(argumentsJson) ?? [:]
|
|
|
|
// 调用工具
|
|
if let result = await mcpClient?.callTool(name: name, arguments: arguments) {
|
|
let functionResult = result["context"] as? String ?? "工具调用失败"
|
|
|
|
// 创建结果对象,格式需要与Android匹配
|
|
let functionCallResult: [String: Any] = ["result": functionResult]
|
|
|
|
// 发送函数调用结果回调
|
|
callback.onFunctionCallResult(functionCall, functionCallResult)
|
|
|
|
// 发送函数调用结果
|
|
sendFunctionCallResult(
|
|
messages: messages,
|
|
systemPrompt: systemPrompt,
|
|
functionCall: functionCall,
|
|
functionResult: functionResult,
|
|
callback: callback
|
|
)
|
|
} else {
|
|
callback.onError(OpenAIError("工具调用失败"))
|
|
}
|
|
} catch {
|
|
print("\(TAG) 自动处理工具调用失败: \(error.localizedDescription)")
|
|
|
|
// 失败时返回错误给回调函数
|
|
let errorMessage = "工具调用失败: \(error.localizedDescription)"
|
|
|
|
// 创建错误结果对象
|
|
let errorResult: [String: Any] = ["error": errorMessage]
|
|
|
|
// 发送函数调用结果
|
|
callback.onFunctionCallResult(functionCall, errorResult)
|
|
|
|
sendFunctionCallResult(
|
|
messages: messages,
|
|
systemPrompt: systemPrompt,
|
|
functionCall: functionCall,
|
|
functionResult: errorMessage,
|
|
callback: callback
|
|
)
|
|
}
|
|
}
|
|
}
|
|
|
|
/// 发送消息(非流式输出)
|
|
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
|
|
}
|
|
}
|
|
|
|
/// 发送消息(流式输出)
|
|
internal func sendMessageStream(
|
|
messages: [[String: Any]],
|
|
systemPrompt: String = "",
|
|
callback: StreamCallback
|
|
) {
|
|
guard isInitialized, !apiKey.isEmpty else {
|
|
callback.onError(OpenAIError("OpenAI服务未初始化"))
|
|
return
|
|
}
|
|
|
|
// 重置取消状态
|
|
isCanceled = false
|
|
|
|
// 检查最后一条消息内容中是否包含图片,决定使用哪个模型
|
|
var currentModel = model
|
|
if let lastMessage = messages.last,
|
|
let content = lastMessage["content"] as? [[String: Any]],
|
|
content.contains(where: { ($0["type"] as? String) == "image_url" }) {
|
|
currentModel = visionModel
|
|
}
|
|
|
|
// 构建请求体
|
|
var requestBody: [String: Any] = [
|
|
"model": currentModel,
|
|
"temperature": 0.7,
|
|
"max_tokens": 2000,
|
|
"stream": true
|
|
]
|
|
|
|
// 构建完整消息数组,添加系统提示
|
|
var fullMessages: [[String: Any]] = []
|
|
if !systemPrompt.isEmpty {
|
|
fullMessages.append(["role": "system", "content": systemPrompt])
|
|
}
|
|
fullMessages.append(contentsOf: messages)
|
|
requestBody["messages"] = fullMessages
|
|
|
|
// 添加工具列表
|
|
if let toolMaps = mcpClient?.getToolMaps(), !toolMaps.isEmpty {
|
|
requestBody["tools"] = toolMaps
|
|
} 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,
|
|
autoHandleMcpTools: autoHandleMcpTools,
|
|
autoToolHandler: { [weak self] functionCall in
|
|
// 如果需要自动处理工具调用,调用处理方法
|
|
guard let self = self, self.autoHandleMcpTools else { return }
|
|
self.autoHandleMcpToolCall(
|
|
functionCall: functionCall,
|
|
messages: messages,
|
|
systemPrompt: systemPrompt,
|
|
callback: callback
|
|
)
|
|
}
|
|
)
|
|
|
|
let sessionConfig = URLSessionConfiguration.default
|
|
let session = URLSession(configuration: sessionConfig, delegate: delegate, delegateQueue: nil)
|
|
let task = session.dataTask(with: request)
|
|
currentStreamTask = task
|
|
task.resume()
|
|
}
|
|
|
|
/// 发送函数调用结果
|
|
internal func sendFunctionCallResult(
|
|
messages: [[String: Any]],
|
|
systemPrompt: String,
|
|
functionCall: [String: Any],
|
|
functionResult: String,
|
|
callback: StreamCallback
|
|
) {
|
|
do {
|
|
// 构建完整消息数组
|
|
var fullMessages: [[String: Any]] = []
|
|
|
|
// 添加系统提示
|
|
if !systemPrompt.isEmpty {
|
|
fullMessages.append(["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": nil as Any?,
|
|
"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
|
|
let autoHandleMcpTools: Bool
|
|
let autoToolHandler: ([String: Any]) -> Void
|
|
|
|
private var buffer = Data()
|
|
private var finalToolCalls: [Int: ToolCallInfo] = [:]
|
|
|
|
init(
|
|
callback: StreamCallback,
|
|
autoHandleMcpTools: Bool = true,
|
|
autoToolHandler: @escaping ([String: Any]) -> Void
|
|
) {
|
|
self.callback = callback
|
|
self.autoHandleMcpTools = autoHandleMcpTools
|
|
self.autoToolHandler = autoToolHandler
|
|
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 {
|
|
print("解析JSON出错: \(error.localizedDescription)")
|
|
// 忽略解析错误,继续处理其他行
|
|
}
|
|
}
|
|
}
|
|
|
|
// 处理工具调用
|
|
private func processToolCalls() -> Bool {
|
|
if finalToolCalls.isEmpty { return false }
|
|
|
|
// 只处理第一个工具调用
|
|
for (_, toolCall) in finalToolCalls {
|
|
if toolCall.isValid {
|
|
// 创建函数调用字典
|
|
let functionCall: [String: Any] = [
|
|
"name": toolCall.name,
|
|
"arguments": toolCall.arguments,
|
|
"id": toolCall.id
|
|
]
|
|
|
|
// 调用回调
|
|
callback.onFunctionCall(functionCall)
|
|
|
|
// 如果需要自动处理工具调用
|
|
if autoHandleMcpTools {
|
|
autoToolHandler(functionCall)
|
|
}
|
|
|
|
return true
|
|
}
|
|
}
|
|
|
|
return false
|
|
}
|
|
|
|
// 处理完成
|
|
func urlSession(_ session: URLSession, task: URLSessionTask, didCompleteWithError error: Error?) {
|
|
if let error = error {
|
|
if (error as NSError).code == NSURLErrorCancelled {
|
|
// 请求被取消,无需处理
|
|
return
|
|
}
|
|
|
|
callback.onError(error)
|
|
} else {
|
|
// 如果没有处理任何工具调用且没有错误,则完成
|
|
if finalToolCalls.isEmpty {
|
|
callback.onComplete()
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|