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.

816 lines
30 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服务流回调类型定义
public typealias StreamCallback = (
onToken: (String) -> Void,
onComplete: () -> Void,
onError: (Error) -> Void,
onFunctionCall: ([String: Any]) -> Void,
onFunctionCallResult: ([String: Any], [String: Any]) -> Void
)
/// 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 isMcpInitialized() -> Bool {
return isMcpInitialized && mcpClient?.isConnected() == 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: @escaping 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
}
}
/// 发送消息(流式输出)
public func sendMessageStream(
messages: [[String: Any]],
systemPrompt: String = "",
callback: @escaping 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()
}
/// 发送函数调用结果
public func sendFunctionCallResult(
messages: [[String: Any]],
systemPrompt: String,
functionCall: [String: Any],
functionResult: String,
callback: @escaping 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: @escaping 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()
}
}
}
}
}