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.

869 lines
30 KiB

import Foundation
import UIKit
import os
import os.log
/// OpenAI服务异常
public struct OpenAIException: Error {
let message: String
public init(_ message: String) {
self.message = message
}
public var localizedDescription: String {
return "OpenAIException: \(message)"
}
}
/// 流式回调协议
public protocol StreamCallback {
func onToken(_ token: String)
func onComplete()
func onError(_ error: Error)
func onFunctionCall(_ functionCall: [String: Any])
func onFunctionCallResult(_ functionCall: [String: Any], _ functionCallResult: [String: Any])
}
/// 工具调用信息
private class ToolCallInfo {
var id: String = ""
var name: String = ""
var arguments: String = ""
func isValid() -> Bool {
return !id.isEmpty && !name.isEmpty
}
}
/// OpenAI服务的原生实现
public class OpenAIService: NSObject {
private let tag = "OpenAIService"
// 日志对象
private let logger = OSLog(subsystem: "com.yunqiinnovation.open_ai_service", category: "OpenAIService")
// MARK: - 属性
private var baseUrl = ""
private var apiKey = ""
private var model = ""
private var visionModel = "doubao-1-5-vision-pro-32k-250115"
private var isInitialized = false
// MCP客户端
private var mcpClient: MCPClient?
private var isMcpInitializedFlag = false
// 流式请求相关
private var currentStreamTask: URLSessionDataTask?
private var currentCallback: StreamCallback?
private var currentMessages: [[String: Any]] = []
private var responseBuffer = Data()
private var toolCalls: [Int: ToolCallInfo] = [:]
private var isCanceled = false
// 优化:添加行缓冲区,避免频繁字符串转换
private var lineBuffer = Data()
private let newlineData = "\n".data(using: .utf8)!
// 优化:添加重连机制
private var retryCount = 0
private let maxRetryCount = 3
private var retryDelay: TimeInterval = 1.0
// 优化:添加性能监控
private var streamStartTime: Date?
private var tokenCount = 0
private var bytesReceived = 0
// URL会话
private lazy var urlSession: URLSession = {
let config = URLSessionConfiguration.default
config.timeoutIntervalForRequest = 30
config.timeoutIntervalForResource = 30
return URLSession(configuration: config)
}()
// 流式会话(使用 delegate)
private lazy var streamSession: URLSession = {
let config = URLSessionConfiguration.default
config.timeoutIntervalForRequest = 30
config.timeoutIntervalForResource = 30
// 优化:减少网络缓存,避免内存压力
config.urlCache = nil
config.requestCachePolicy = .reloadIgnoringLocalCacheData
// 优化:限制并发连接数
config.httpMaximumConnectionsPerHost = 2
// 优化:针对后台数据处理的线程配置
let delegateQueue = OperationQueue()
// 使用utility QoS,平衡性能和功耗,不与UI线程竞争
delegateQueue.qualityOfService = .utility
delegateQueue.maxConcurrentOperationCount = 1
delegateQueue.name = "OpenAIService.StreamDelegate"
return URLSession(configuration: config, delegate: self, delegateQueue: delegateQueue)
}()
// MARK: - 初始化
public override init() {
super.init()
}
// MARK: - 公共方法
/// 初始化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 mcpClient == nil {
mcpClient = MCPClient()
}
_ = initializeMcpClient(serverUrl: mcpServer)
isInitialized = !apiKey.isEmpty
return isInitialized
}
/// 创建用户消息
public func createUserMessage(content: String) -> [String: Any] {
return [
"role": "user",
"content": content
]
}
/// 创建助手消息
public func createAssistantMessage(content: String) -> [String: Any] {
return [
"role": "assistant",
"content": content
]
}
/// 创建系统消息
public func createSystemMessage(content: String) -> [String: Any] {
return [
"role": "system",
"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
]
}
/// 发送消息(非流式输出)
public func sendMessage(messages: [[String: Any]]) throws -> String {
guard isInitialized && !apiKey.isEmpty else {
throw OpenAIException("OpenAI服务未初始化")
}
// 检查是否包含图片,决定使用哪个模型
var currentModel = model
if let lastMessage = messages.last,
let content = lastMessage["content"] as? String,
content.contains("image_url") {
currentModel = visionModel
}
// 构建请求体
var requestBody: [String: Any] = [
"model": currentModel,
"messages": messages,
"temperature": 0.7,
"max_tokens": 2000,
"stream": false
]
// 添加工具列表
if let tools = mcpClient?.getToolMaps(), !tools.isEmpty {
requestBody["tools"] = tools
}
// 创建请求
guard let url = URL(string: baseUrl),
let jsonData = try? JSONSerialization.data(withJSONObject: requestBody) else {
throw OpenAIException("创建请求失败")
}
var request = URLRequest(url: url)
request.httpMethod = "POST"
request.setValue("application/json", forHTTPHeaderField: "Content-Type")
request.setValue("Bearer \(apiKey)", forHTTPHeaderField: "Authorization")
request.httpBody = jsonData
// 同步请求
let semaphore = DispatchSemaphore(value: 0)
var result: String = ""
var error: Error?
let task = urlSession.dataTask(with: request) { data, response, taskError in
defer { semaphore.signal() }
if let taskError = taskError {
error = OpenAIException("请求失败: \(taskError.localizedDescription)")
return
}
guard let httpResponse = response as? HTTPURLResponse else {
error = OpenAIException("无效响应")
return
}
guard httpResponse.statusCode == 200 else {
error = OpenAIException("API调用失败: \(httpResponse.statusCode)")
return
}
guard let data = data,
let jsonResponse = try? JSONSerialization.jsonObject(with: data) as? [String: Any] else {
error = OpenAIException("响应解析失败")
return
}
// 检查是否有函数调用
if let choices = jsonResponse["choices"] as? [[String: Any]],
let choice = choices.first,
let message = choice["message"] as? [String: Any] {
// 检查是否有工具调用
if let toolCalls = message["tool_calls"] as? [[String: Any]],
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 functionCallJson: [String: Any] = [
"name": name,
"arguments": arguments,
"id": id
]
if let jsonData = try? JSONSerialization.data(withJSONObject: functionCallJson),
let jsonString = String(data: jsonData, encoding: .utf8) {
result = jsonString
} else {
error = OpenAIException("函数调用序列化失败")
}
return
}
// 如果没有工具调用,返回消息内容
if let content = message["content"] as? String {
result = content
return
}
}
error = OpenAIException("无效的响应格式")
}
task.resume()
semaphore.wait()
if let error = error {
throw error
}
return result
}
/// 发送消息(流式输出)
public func sendMessageStream(messages: [[String: Any]], callback: StreamCallback) {
guard isInitialized && !apiKey.isEmpty else {
callback.onError(OpenAIException("OpenAI服务未初始化"))
return
}
// 重置状态
isCanceled = false
currentCallback = callback
currentMessages = messages
responseBuffer = Data()
toolCalls = [:]
// 优化:初始化性能监控
streamStartTime = Date()
tokenCount = 0
bytesReceived = 0
// 检查是否包含图片,决定使用哪个模型
var currentModel = model
if let lastMessage = messages.last,
let content = lastMessage["content"] as? String,
content.contains("image_url") {
currentModel = visionModel
}
// 构建请求体
var requestBody: [String: Any] = [
"model": currentModel,
"messages": messages,
"temperature": 0.7,
"max_tokens": 2000,
"stream": true
]
// 添加工具列表
if let tools = mcpClient?.getToolMaps(), !tools.isEmpty {
requestBody["tools"] = tools
}
// 创建请求
guard let url = URL(string: baseUrl),
let jsonData = try? JSONSerialization.data(withJSONObject: requestBody) else {
callback.onError(OpenAIException("创建请求失败"))
return
}
var request = URLRequest(url: url)
request.httpMethod = "POST"
request.setValue("application/json", forHTTPHeaderField: "Content-Type")
request.setValue("Bearer \(apiKey)", forHTTPHeaderField: "Authorization")
request.setValue("text/event-stream", forHTTPHeaderField: "Accept")
request.setValue("cache-control", forHTTPHeaderField: "no-cache")
request.httpBody = jsonData
os_log("开始流式请求: %{public}@", log: logger, type: .info, url.absoluteString)
// 使用 delegate 模式的 URLSession 创建任务
let task = streamSession.dataTask(with: request)
currentStreamTask = task
task.resume()
}
/// 处理缓冲区数据
private func processBufferedData(callback: StreamCallback, isComplete: Bool = false) {
// 优化:使用流式解析,避免频繁的字符串转换
processStreamData(callback: callback, isComplete: isComplete)
}
/// 优化的流式数据处理
private func processStreamData(callback: StreamCallback, isComplete: Bool = false) {
lineBuffer.append(responseBuffer)
responseBuffer = Data()
var searchRange = lineBuffer.startIndex..<lineBuffer.endIndex
while let newlineRange = lineBuffer.range(of: newlineData, in: searchRange) {
// 提取完整行
let lineData = lineBuffer.subdata(in: lineBuffer.startIndex..<newlineRange.lowerBound)
// 转换为字符串并处理
if let lineString = String(data: lineData, encoding: .utf8) {
processStreamLine(lineString.trimmingCharacters(in: .whitespacesAndNewlines), callback: callback)
}
// 移除已处理的行
lineBuffer.removeSubrange(lineBuffer.startIndex..<newlineRange.upperBound)
searchRange = lineBuffer.startIndex..<lineBuffer.endIndex
}
// 如果是完成状态,处理剩余数据
if isComplete && !lineBuffer.isEmpty {
if let lineString = String(data: lineBuffer, encoding: .utf8) {
processStreamLine(lineString.trimmingCharacters(in: .whitespacesAndNewlines), callback: callback)
}
lineBuffer = Data()
}
}
/// 处理单行流式数据
private func processStreamLine(_ line: String, callback: StreamCallback) {
guard !line.isEmpty else { return }
if line.hasPrefix("data:") {
let dataContent = String(line.dropFirst(5)).trimmingCharacters(in: .whitespacesAndNewlines)
// 处理[DONE]消息
if dataContent == "[DONE]" || dataContent == "[\"DONE\"]" {
os_log("收到完成信号", log: logger, type: .info)
let hasToolCalls = processToolCalls(callback: callback)
if !hasToolCalls {
// 直接回调,由调用方处理线程切换
callback.onComplete()
}
return
}
// 优化:直接尝试解析JSON,避免预验证
if let jsonData = dataContent.data(using: .utf8) {
do {
if let json = try JSONSerialization.jsonObject(with: jsonData) as? [String: Any] {
handleStreamJson(json: json, callback: callback)
}
} catch {
// 静默忽略无效JSON,避免日志噪音
}
}
}
}
/// 处理流式JSON数据
private func handleStreamJson(json: [String: Any], callback: StreamCallback) {
guard let choices = json["choices"] as? [[String: Any]],
let choice = choices.first,
let delta = choice["delta"] as? [String: Any] else {
return
}
// 处理普通文本内容
if let content = delta["content"] as? String {
if !isCanceled {
// 优化:更新token计数
tokenCount += 1
// 直接回调,由调用方处理线程切换
callback.onToken(content)
}
}
// 收集工具调用信息
if let toolCallsArray = delta["tool_calls"] as? [[String: Any]] {
for toolCallDict in toolCallsArray {
guard let index = toolCallDict["index"] as? Int else { continue }
// 创建或获取现有的工具调用信息
if toolCalls[index] == nil {
toolCalls[index] = ToolCallInfo()
}
let toolCallInfo = toolCalls[index]!
// 更新ID
if let id = toolCallDict["id"] as? String {
toolCallInfo.id = id
}
// 更新函数信息
if let function = toolCallDict["function"] as? [String: Any] {
if let name = function["name"] as? String {
toolCallInfo.name = name
}
if let arguments = function["arguments"] as? String {
toolCallInfo.arguments += arguments
}
}
}
}
}
/// 处理工具调用
private func processToolCalls(callback: StreamCallback) -> Bool {
guard let firstToolCall = toolCalls.values.first, firstToolCall.isValid() else {
return false
}
// 创建函数调用字典
let functionCall: [String: Any] = [
"name": firstToolCall.name,
"arguments": firstToolCall.arguments,
"id": firstToolCall.id
]
os_log("工具调用: id=%{public}@, name=%{public}@, arguments=%{public}@", log: logger, type: .info, firstToolCall.id, firstToolCall.name, firstToolCall.arguments)
// 通知上层工具调用事件
callback.onFunctionCall(functionCall)
// 在后台队列处理工具调用
Task {
do {
if !self.isCanceled {
// 调用工具
let result: [String: Any]
do {
let args = try self.parseJsonArguments(firstToolCall.arguments)
result = try await self.mcpClient?.callTool(name: firstToolCall.name, arguments: args) ?? ["context": "工具调用失败"]
} catch {
os_log("工具调用错误: %{public}@", log: logger, type: .error, error.localizedDescription)
result = ["context": "工具调用失败: \(error.localizedDescription)"]
}
if !self.isCanceled {
// 处理结果
callback.onFunctionCallResult(functionCall, result)
let context = result["context"] as? String ?? "工具调用失败"
os_log("MCP工具调用完成: name=%{public}@, context=%{public}@", log: logger, type: .info, firstToolCall.name, context)
// 将结果发送回OpenAI继续对话
self.sendFunctionCallResult(
messages: self.currentMessages,
functionCall: functionCall,
functionResult: context,
callback: callback
)
}
}
} catch {
if !self.isCanceled {
os_log("处理工具调用异常: %{public}@", log: logger, type: .error, error.localizedDescription)
let errorMessage = "工具调用处理失败: \(error.localizedDescription)"
self.sendFunctionCallResult(
messages: self.currentMessages,
functionCall: functionCall,
functionResult: errorMessage,
callback: callback
)
}
}
}
return true
}
/// 发送函数调用结果
public func sendFunctionCallResult(messages: [[String: Any]], functionCall: [String: Any], functionResult: String, callback: StreamCallback) {
if isCanceled {
os_log("请求已取消,不发送函数调用结果", log: logger, type: .info)
return
}
var fullMessages = messages
// 添加函数调用消息
let callId = functionCall["id"] as? String ?? "call_\(Int(Date().timeIntervalSince1970))"
fullMessages.append([
"role": "assistant",
"content": "",
"tool_calls": [[
"id": callId,
"type": "function",
"function": [
"name": functionCall["name"] as? String ?? "",
"arguments": functionCall["arguments"] as? String ?? ""
]
]]
])
// 添加函数调用结果
fullMessages.append([
"role": "tool",
"content": functionResult,
"tool_call_id": callId
])
// 发送完整对话
sendMessageStream(messages: fullMessages, callback: callback)
}
/// 取消当前流式请求
public func cancelCurrentStream() -> Bool {
isCanceled = true
currentStreamTask?.cancel()
// 优化:更彻底的状态清理
cleanupStreamState()
os_log("已取消当前流式请求", log: logger, type: .info)
return true
}
/// 优化:统一的状态清理方法
private func cleanupStreamState() {
currentStreamTask = nil
currentCallback = nil
currentMessages.removeAll()
responseBuffer = Data()
lineBuffer = Data()
toolCalls.removeAll()
retryCount = 0
// 优化:重置性能监控变量
streamStartTime = nil
tokenCount = 0
bytesReceived = 0
}
/// 取消所有操作并释放资源
public func cancelAll() {
_ = cancelCurrentStream()
}
// MARK: - MCP相关方法
/// 初始化MCP客户端
public func initializeMcpClient(serverUrl: String) -> Bool {
if mcpClient != nil {
mcpClient?.close()
}
mcpClient = MCPClient()
Task {
do {
let result = try await mcpClient?.connectToSSE(mcpConfigJson: serverUrl) ?? false
isMcpInitializedFlag = result
os_log("MCP客户端初始化%{public}@", log: logger, type: .info, result ? "成功" : "失败")
} catch {
os_log("MCP客户端初始化失败: %{public}@", log: logger, type: .error, error.localizedDescription)
isMcpInitializedFlag = false
}
}
return true // 立即返回,实际连接在后台进行
}
/// MCP客户端是否已初始化
public func isMcpInitialized() -> Bool {
return isMcpInitializedFlag && mcpClient?.isConnected() == true
}
/// 关闭MCP客户端
public func closeMcpClient() {
mcpClient?.close()
mcpClient = nil
isMcpInitializedFlag = false
}
/// 注册函数
public func registerFunction(name: String, description: String, parameters: [String: Any]) -> Bool {
guard let mcpClient = mcpClient else {
os_log("MCP客户端未初始化", log: logger, type: .error)
return false
}
// 创建函数处理器
let handler = LocalFunctionHandler(name: name)
// 注册本地函数
return mcpClient.registerLocalFunction(
name: name,
description: description,
parameters: parameters,
handler: handler
)
}
/// 处理MCP工具调用
public func handleMcpToolCall(functionCallJson: String) async throws -> String {
guard let mcpClient = mcpClient, isMcpInitializedFlag else {
return "MCP客户端未初始化"
}
guard let jsonData = functionCallJson.data(using: .utf8),
let functionCall = try? JSONSerialization.jsonObject(with: jsonData) as? [String: Any],
let name = functionCall["name"] as? String,
let argumentsJson = functionCall["arguments"] as? String else {
throw OpenAIException("解析函数调用失败")
}
let arguments = try parseJsonArguments(argumentsJson)
let result = try await mcpClient.callTool(name: name, arguments: arguments)
return result["context"] as? String ?? "无法处理MCP工具调用"
}
// MARK: - 工具方法
/// 解析JSON参数
private func parseJsonArguments(_ argumentsJson: String) throws -> [String: Any] {
guard let data = argumentsJson.data(using: .utf8),
let arguments = try? JSONSerialization.jsonObject(with: data) as? [String: Any] else {
return [:]
}
return arguments
}
/// 检查字符串是否为有效的JSON
private func isValidJson(jsonString: String) -> Bool {
guard !jsonString.isEmpty,
jsonString.hasPrefix("{"),
jsonString.hasSuffix("}") else {
return false
}
var braceCount = 0
var insideQuotes = false
var escapeNext = false
for char in jsonString {
if escapeNext {
escapeNext = false
continue
}
switch char {
case "\\":
escapeNext = true
case "\"":
if !escapeNext {
insideQuotes.toggle()
}
case "{":
if !insideQuotes {
braceCount += 1
}
case "}":
if !insideQuotes {
braceCount -= 1
}
default:
break
}
// 如果括号不匹配(闭合太多),提前返回
if braceCount < 0 {
return false
}
}
// 所有引号都应该闭合,所有括号也应该匹配
return braceCount == 0 && !insideQuotes
}
/// 将文件转换为Base64字符串
public func fileToBase64(filePath: String, maxSizeKB: Int = 20480) -> String? {
guard let image = UIImage(contentsOfFile: filePath) else {
os_log("文件不存在或无法读取: %{public}@", log: logger, type: .error, filePath)
return nil
}
var processedImage = image
// 检查图片尺寸,限制最大为1024*1024
let maxDimension: CGFloat = 1024
if image.size.width > maxDimension || image.size.height > maxDimension {
let ratio = min(maxDimension / image.size.width, maxDimension / image.size.height)
let newSize = CGSize(width: image.size.width * ratio, height: image.size.height * ratio)
UIGraphicsBeginImageContextWithOptions(newSize, false, 1.0)
image.draw(in: CGRect(origin: .zero, size: newSize))
processedImage = UIGraphicsGetImageFromCurrentImageContext() ?? image
UIGraphicsEndImageContext()
}
// 压缩图片
guard let imageData = processedImage.jpegData(compressionQuality: 0.8) else {
return nil
}
return imageData.base64EncodedString()
}
}
// MARK: - URLSessionDataDelegate Extension
extension OpenAIService: URLSessionDataDelegate {
/// 接收到响应头
public func urlSession(_ session: URLSession, dataTask: URLSessionDataTask, didReceive response: URLResponse, completionHandler: @escaping (URLSession.ResponseDisposition) -> Void) {
guard let httpResponse = response as? HTTPURLResponse else {
if !isCanceled {
// 直接回调,由调用方处理线程切换
currentCallback?.onError(OpenAIException("无效响应"))
}
completionHandler(.cancel)
return
}
guard httpResponse.statusCode == 200 else {
if !isCanceled {
// 直接回调,由调用方处理线程切换
currentCallback?.onError(OpenAIException("API调用失败: \(httpResponse.statusCode)"))
}
completionHandler(.cancel)
return
}
os_log("收到响应头,状态码: %d", log: logger, type: .info, httpResponse.statusCode)
completionHandler(.allow)
}
/// 实时接收数据块
public func urlSession(_ session: URLSession, dataTask: URLSessionDataTask, didReceive data: Data) {
guard !isCanceled, let callback = currentCallback else { return }
// 优化:更新性能指标
bytesReceived += data.count
// 将新数据添加到缓冲区
responseBuffer.append(data)
// 处理缓冲区中的完整行
processBufferedData(callback: callback)
}
/// 请求完成
public func urlSession(_ session: URLSession, dataTask: URLSessionDataTask, didCompleteWithError error: Error?) {
defer {
// 优化:使用统一的清理方法
cleanupStreamState()
}
guard !isCanceled else {
os_log("流式请求已取消", log: logger, type: .info)
return
}
if let error = error {
os_log("流式请求错误: %{public}@", log: logger, type: .error, error.localizedDescription)
// 直接回调,由调用方处理线程切换
currentCallback?.onError(OpenAIException("请求失败: \(error.localizedDescription)"))
return
}
// 处理剩余的缓冲区数据
if let callback = currentCallback {
processBufferedData(callback: callback, isComplete: true)
}
// 优化:输出性能统计
if let startTime = streamStartTime {
let duration = Date().timeIntervalSince(startTime)
let tokensPerSecond = duration > 0 ? Double(tokenCount) / duration : 0
os_log("流式请求性能统计 - 耗时: %.2fs, tokens: %d, 字节: %d, tokens/s: %.2f",
log: logger, type: .info, duration, tokenCount, bytesReceived, tokensPerSecond)
}
os_log("流式请求完成", log: logger, type: .info)
}
}
/// 本地函数处理器
private class LocalFunctionHandler: FunctionHandler {
private let functionName: String
init(name: String) {
self.functionName = name
}
func handle(arguments: [String: Any]) async throws -> String {
// 由于本地函数的实际处理是在Flutter端完成的
// 这里只需返回一个标记,表示该函数是本地函数
return "LOCAL_FUNCTION:\(functionName)"
}
}