22 changed files with 2465 additions and 444 deletions
Binary file not shown.
Binary file not shown.
Binary file not shown.
@ -0,0 +1,135 @@ |
|||
# OpenAI API 与 MCP 集成插件 |
|||
|
|||
这个 Flutter 插件提供了 OpenAI API 的访问能力和 Model Context Protocol (MCP) 工具调用功能的集成。 |
|||
|
|||
## 功能特点 |
|||
|
|||
- 使用官方 OpenAI Java SDK 进行异步通信 |
|||
- 集成官方 MCP Kotlin SDK,使用 SSE 模式 |
|||
- 支持流式输出响应 |
|||
- 支持工具调用和结果处理 |
|||
- 支持带图片的多模态对话 |
|||
|
|||
## MCP 集成 |
|||
|
|||
本插件使用 [MCP 官方 Kotlin SDK](https://github.com/modelcontextprotocol/kotlin-sdk) 实现与 MCP 服务器的通信。通过 SSE(Server-Sent Events)模式连接,能够: |
|||
|
|||
- 获取 MCP 服务器提供的所有工具定义 |
|||
- 动态调用远程工具并获取结果 |
|||
- 支持本地工具的注册和调用 |
|||
- 在 OpenAI API 请求中无缝集成工具调用功能 |
|||
|
|||
## 安装 |
|||
|
|||
在项目的 `pubspec.yaml` 中添加本地插件依赖: |
|||
|
|||
```yaml |
|||
dependencies: |
|||
open_ai: |
|||
path: local_plugins/open_ai |
|||
``` |
|||
|
|||
## 使用方法 |
|||
|
|||
### 初始化 |
|||
|
|||
```dart |
|||
import 'package:open_ai/open_ai.dart'; |
|||
|
|||
final openAI = OpenAI(); |
|||
|
|||
await openAI.initialize( |
|||
apiKey: 'your-api-key', |
|||
baseUrl: 'https://api.example.com', // 可选,默认为 OpenAI 官方 API |
|||
model: 'gpt-3.5-turbo', // 可选,默认为 gpt-3.5-turbo |
|||
mcpServer: 'https://mcp.example.com', // MCP 服务器 SSE 端点地址 |
|||
); |
|||
``` |
|||
|
|||
### 创建消息 |
|||
|
|||
```dart |
|||
// 创建系统消息 |
|||
final systemMessage = await openAI.createSystemMessage('你是一个助手'); |
|||
|
|||
// 创建用户消息 |
|||
final userMessage = await openAI.createUserMessage('你好,请帮我解释一下量子力学'); |
|||
|
|||
// 创建助手消息 |
|||
final assistantMessage = await openAI.createAssistantMessage('我可以帮你解释量子力学'); |
|||
|
|||
// 创建带图片的用户消息 |
|||
final imageBase64 = '...'; // base64编码的图片数据 |
|||
final userImageMessage = await openAI.createUserMessageWithImage( |
|||
'这张图片中的物体是什么?', |
|||
imageBase64 |
|||
); |
|||
``` |
|||
|
|||
### 发送消息(非流式输出) |
|||
|
|||
```dart |
|||
final messages = [systemMessage, userMessage]; |
|||
final response = await openAI.sendMessage(messages); |
|||
print('AI回复: $response'); |
|||
``` |
|||
|
|||
### 发送消息(流式输出) |
|||
|
|||
```dart |
|||
final messages = [systemMessage, userMessage]; |
|||
final callback = StreamCallback( |
|||
onToken: (token) { |
|||
// 处理单个令牌 |
|||
print('收到令牌: $token'); |
|||
}, |
|||
onComplete: () { |
|||
// 处理完成事件 |
|||
print('响应完成'); |
|||
}, |
|||
onError: (error) { |
|||
// 处理错误 |
|||
print('发生错误: $error'); |
|||
}, |
|||
onFunctionCall: (functionCall) { |
|||
// 处理函数调用 |
|||
print('函数调用: $functionCall'); |
|||
}, |
|||
onFunctionCallResult: (functionCall, result) { |
|||
// 处理函数调用结果 |
|||
print('函数调用结果: $result'); |
|||
}, |
|||
); |
|||
|
|||
final streamId = await openAI.sendMessageStream(messages, callback); |
|||
``` |
|||
|
|||
### 取消当前流式请求 |
|||
|
|||
```dart |
|||
final success = await openAI.cancelCurrentStream(); |
|||
``` |
|||
|
|||
### 释放资源 |
|||
|
|||
```dart |
|||
await openAI.dispose(); |
|||
``` |
|||
|
|||
## 异常处理 |
|||
|
|||
该插件会在操作失败时抛出异常,请使用 try-catch 块捕获它们: |
|||
|
|||
```dart |
|||
try { |
|||
final response = await openAI.sendMessage(messages); |
|||
} catch (e) { |
|||
print('发生错误: $e'); |
|||
} |
|||
``` |
|||
|
|||
## 注意事项 |
|||
|
|||
- 初始化插件时必须提供有效的 API 密钥 |
|||
- 使用流式响应时,请确保在完成后调用 `dispose()` 方法释放资源 |
|||
- MCP 功能需要有效的 MCP 服务器 SSE 端点地址才能工作 |
|||
@ -0,0 +1,71 @@ |
|||
group = "com.yunqiinnovation.open_ai" |
|||
version = "1.0" |
|||
|
|||
buildscript { |
|||
val kotlinVersion by extra("1.8.0") |
|||
repositories { |
|||
google() |
|||
mavenCentral() |
|||
} |
|||
|
|||
dependencies { |
|||
classpath("com.android.tools.build:gradle:7.3.0") |
|||
classpath("org.jetbrains.kotlin:kotlin-gradle-plugin:$kotlinVersion") |
|||
} |
|||
} |
|||
|
|||
allprojects { |
|||
repositories { |
|||
google() |
|||
mavenCentral() |
|||
} |
|||
} |
|||
|
|||
plugins { |
|||
id("com.android.library") |
|||
id("kotlin-android") |
|||
} |
|||
|
|||
android { |
|||
compileSdk = 33 |
|||
|
|||
compileOptions { |
|||
sourceCompatibility = JavaVersion.VERSION_1_8 |
|||
targetCompatibility = JavaVersion.VERSION_1_8 |
|||
} |
|||
|
|||
kotlinOptions { |
|||
jvmTarget = "1.8" |
|||
} |
|||
|
|||
defaultConfig { |
|||
minSdk = 21 |
|||
} |
|||
|
|||
namespace = "com.yunqiinnovation.open_ai" |
|||
} |
|||
|
|||
dependencies { |
|||
val kotlinVersion: String by project |
|||
|
|||
implementation("org.jetbrains.kotlin:kotlin-stdlib-jdk7:$kotlinVersion") |
|||
implementation("androidx.annotation:annotation:1.6.0") |
|||
implementation("androidx.appcompat:appcompat:1.6.1") |
|||
|
|||
// Kotlin协程 |
|||
implementation("org.jetbrains.kotlinx:kotlinx-coroutines-android:1.7.1") |
|||
|
|||
// OkHttp 依赖项 |
|||
implementation("com.squareup.okhttp3:okhttp:4.11.0") |
|||
implementation("com.squareup.okhttp3:okhttp-sse:4.11.0") |
|||
|
|||
// Jackson JSON解析器 |
|||
implementation("com.fasterxml.jackson.core:jackson-databind:2.14.2") |
|||
|
|||
// OpenAI Java库依赖 |
|||
implementation("com.aallam.openai:openai-client:3.6.0") |
|||
implementation("io.ktor:ktor-client-okhttp:2.3.3") |
|||
|
|||
// MCP官方Kotlin SDK依赖 |
|||
implementation("io.modelcontextprotocol:kotlin-sdk:0.5.0") |
|||
} |
|||
@ -0,0 +1,258 @@ |
|||
package com.yunqiinnovation.open_ai |
|||
|
|||
import android.content.Context |
|||
import android.util.Log |
|||
import com.fasterxml.jackson.databind.ObjectMapper |
|||
import io.modelcontextprotocol.kotlin.sdk.Implementation |
|||
import io.modelcontextprotocol.kotlin.sdk.client.Client |
|||
import io.modelcontextprotocol.kotlin.sdk.client.SseClientTransport |
|||
import io.modelcontextprotocol.kotlin.sdk.tools.ToolCallRequest |
|||
import io.modelcontextprotocol.kotlin.sdk.tools.ToolDefinition |
|||
import kotlinx.coroutines.* |
|||
import kotlinx.coroutines.channels.Channel |
|||
import kotlinx.coroutines.channels.awaitClose |
|||
import kotlinx.coroutines.flow.Flow |
|||
import kotlinx.coroutines.flow.callbackFlow |
|||
import okhttp3.* |
|||
import okhttp3.sse.EventSource |
|||
import okhttp3.sse.EventSourceListener |
|||
import okhttp3.sse.EventSources |
|||
import org.json.JSONArray |
|||
import org.json.JSONObject |
|||
import java.io.IOException |
|||
import java.util.concurrent.TimeUnit |
|||
import kotlin.coroutines.CoroutineContext |
|||
|
|||
/** |
|||
* MCP功能处理接口 |
|||
*/ |
|||
interface FunctionHandler { |
|||
suspend fun handle(arguments: Map<String, Any>): String |
|||
} |
|||
|
|||
/** |
|||
* Model Context Protocol (MCP)客户端实现类 |
|||
* 使用官方 MCP Kotlin SDK |
|||
*/ |
|||
class MCPClient(private val context: Context? = null) : CoroutineScope { |
|||
private val TAG = "MCPClient" |
|||
|
|||
// 协程相关 |
|||
private val job = SupervisorJob() |
|||
override val coroutineContext: CoroutineContext |
|||
get() = Dispatchers.IO + job |
|||
|
|||
// OkHttp客户端 |
|||
private val httpClient = OkHttpClient.Builder() |
|||
.connectTimeout(30, TimeUnit.SECONDS) |
|||
.readTimeout(30, TimeUnit.SECONDS) |
|||
.writeTimeout(30, TimeUnit.SECONDS) |
|||
.build() |
|||
|
|||
// MCP客户端 |
|||
private var client: Client? = null |
|||
|
|||
// 已注册的工具 |
|||
private val tools = mutableListOf<String>() |
|||
|
|||
// 本地函数处理器 |
|||
private val functionHandlers = mutableMapOf<String, FunctionHandler>() |
|||
|
|||
// Json解析器 |
|||
private val objectMapper = ObjectMapper() |
|||
|
|||
// 连接状态 |
|||
private var isConnected = false |
|||
|
|||
/** |
|||
* 使用SSE方式连接到MCP服务器 |
|||
*/ |
|||
suspend fun connectToSSE(serverUrl: String): Boolean { |
|||
if (serverUrl.isEmpty()) { |
|||
Log.e(TAG, "服务器URL为空") |
|||
return false |
|||
} |
|||
|
|||
try { |
|||
// 创建MCP客户端 |
|||
client = Client( |
|||
clientInfo = Implementation( |
|||
name = "deepvoice-mcp-client", |
|||
version = "1.0.0" |
|||
) |
|||
) |
|||
|
|||
// 创建SSE传输 |
|||
val transport = SseClientTransport( |
|||
serverUrl = serverUrl, |
|||
httpClient = httpClient |
|||
) |
|||
|
|||
// 连接到服务器 |
|||
client?.connect(transport) |
|||
isConnected = true |
|||
|
|||
// 获取工具定义 |
|||
fetchTools() |
|||
|
|||
return true |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "连接到MCP服务器失败: ${e.message}", e) |
|||
isConnected = false |
|||
return false |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 获取工具定义 |
|||
*/ |
|||
private suspend fun fetchTools() { |
|||
try { |
|||
val availableTools = client?.listTools() ?: emptyList() |
|||
tools.clear() |
|||
|
|||
availableTools.forEach { toolDefinition -> |
|||
val toolJson = toolDefinitionToJson(toolDefinition) |
|||
tools.add(toolJson.toString()) |
|||
Log.d(TAG, "已获取工具: ${toolDefinition.name}") |
|||
} |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "获取工具定义失败: ${e.message}", e) |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 将工具定义转换为JSON |
|||
*/ |
|||
private fun toolDefinitionToJson(toolDefinition: ToolDefinition): JSONObject { |
|||
val functionObject = JSONObject().apply { |
|||
put("name", toolDefinition.name) |
|||
put("description", toolDefinition.description ?: "") |
|||
|
|||
// 处理参数定义 |
|||
if (toolDefinition.parameters != null) { |
|||
put("parameters", JSONObject(toolDefinition.parameters)) |
|||
} |
|||
} |
|||
|
|||
return JSONObject().apply { |
|||
put("type", "function") |
|||
put("function", functionObject) |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 获取当前连接状态 |
|||
*/ |
|||
fun isConnected(): Boolean { |
|||
return isConnected && client != null |
|||
} |
|||
|
|||
/** |
|||
* 注册本地函数 |
|||
*/ |
|||
fun registerLocalFunction( |
|||
name: String, |
|||
description: String, |
|||
parameters: JSONObject, |
|||
handler: FunctionHandler |
|||
): Boolean { |
|||
try { |
|||
// 创建工具定义 |
|||
val toolDefinition = JSONObject().apply { |
|||
put("type", "function") |
|||
put("function", JSONObject().apply { |
|||
put("name", name) |
|||
put("description", description) |
|||
put("parameters", parameters) |
|||
}) |
|||
} |
|||
|
|||
// 添加到工具列表 |
|||
tools.add(toolDefinition.toString()) |
|||
|
|||
// 注册处理器 |
|||
functionHandlers[name] = handler |
|||
|
|||
Log.d(TAG, "已注册本地函数: $name") |
|||
return true |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "注册本地函数失败: ${e.message}", e) |
|||
return false |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 获取所有工具的定义 |
|||
*/ |
|||
fun getToolMaps(): List<String> { |
|||
return tools.toList() |
|||
} |
|||
|
|||
/** |
|||
* 解析JSON参数 |
|||
*/ |
|||
fun parseJsonArguments(argumentsJson: String): Map<String, Any> { |
|||
try { |
|||
return objectMapper.readValue(argumentsJson, Map::class.java) as Map<String, Any> |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "解析JSON参数失败: ${e.message}", e) |
|||
return mapOf() |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 调用工具 |
|||
*/ |
|||
suspend fun callTool(name: String, arguments: Map<String, Any>): JSONObject { |
|||
try { |
|||
// 检查是否是本地函数 |
|||
if (functionHandlers.containsKey(name)) { |
|||
val handler = functionHandlers[name] |
|||
val result = handler?.handle(arguments) ?: throw Exception("函数处理器为空") |
|||
|
|||
// 函数处理结果格式化为JSON |
|||
return JSONObject().apply { |
|||
put("context", result) |
|||
} |
|||
} |
|||
|
|||
// 如果不是本地函数,使用MCP客户端调用远程工具 |
|||
if (client != null && isConnected) { |
|||
val toolCallRequest = ToolCallRequest( |
|||
name = name, |
|||
arguments = objectMapper.writeValueAsString(arguments) |
|||
) |
|||
|
|||
val result = client?.callTool(toolCallRequest) |
|||
return JSONObject().apply { |
|||
put("context", result?.result ?: "工具调用失败,未收到结果") |
|||
} |
|||
} |
|||
|
|||
throw Exception("MCP客户端未连接") |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "调用工具失败: ${e.message}", e) |
|||
return JSONObject().apply { |
|||
put("context", "调用工具失败: ${e.message}") |
|||
} |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 关闭客户端 |
|||
*/ |
|||
fun close() { |
|||
try { |
|||
client?.disconnect() |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "关闭MCP客户端失败: ${e.message}", e) |
|||
} finally { |
|||
client = null |
|||
isConnected = false |
|||
job.cancel() // 取消所有协程 |
|||
tools.clear() // 清除工具列表 |
|||
functionHandlers.clear() // 清除函数处理器 |
|||
} |
|||
} |
|||
} |
|||
@ -0,0 +1,988 @@ |
|||
package com.yunqiinnovation.open_ai |
|||
|
|||
import android.content.Context |
|||
import android.graphics.Bitmap |
|||
import android.graphics.BitmapFactory |
|||
import android.util.Base64 |
|||
import android.util.Log |
|||
import com.aallam.openai.api.chat.* |
|||
import com.aallam.openai.api.file.FileSource |
|||
import com.aallam.openai.api.http.Timeout |
|||
import com.aallam.openai.api.image.ImageCreation |
|||
import com.aallam.openai.api.model.ModelId |
|||
import com.aallam.openai.client.OpenAI |
|||
import com.aallam.openai.client.OpenAIConfig |
|||
import com.aallam.openai.client.OpenAIHost |
|||
import com.fasterxml.jackson.databind.ObjectMapper |
|||
import kotlinx.coroutines.* |
|||
import kotlinx.coroutines.flow.* |
|||
import okhttp3.* |
|||
import okhttp3.sse.EventSource |
|||
import okhttp3.sse.EventSourceListener |
|||
import okhttp3.sse.EventSources |
|||
import org.json.JSONArray |
|||
import org.json.JSONObject |
|||
import java.io.ByteArrayOutputStream |
|||
import java.io.File |
|||
import java.io.IOException |
|||
import java.util.UUID |
|||
import java.util.concurrent.TimeUnit |
|||
import kotlin.coroutines.CoroutineContext |
|||
import kotlin.time.Duration.Companion.seconds |
|||
|
|||
/** |
|||
* OpenAI服务的原生实现 |
|||
* 集成了OpenAI官方Java SDK和MCP的官方Kotlin SDK |
|||
*/ |
|||
class OpenAIService(private val context: Context? = null) : CoroutineScope { |
|||
private val TAG = "OpenAIService" |
|||
|
|||
// 协程相关 |
|||
private val job = SupervisorJob() |
|||
override val coroutineContext: CoroutineContext |
|||
get() = Dispatchers.IO + job |
|||
|
|||
// 添加辅助方法,确保回调在主线程执行 |
|||
private suspend fun safeCallback(block: suspend () -> Unit) { |
|||
withContext(Dispatchers.Main) { |
|||
block() |
|||
} |
|||
} |
|||
|
|||
// OpenAI客户端 |
|||
private var openAI: OpenAI? = null |
|||
private var baseUrl = "https://api.openai.com/v1/chat/completions" |
|||
private var apiKey: String = "" |
|||
private var isInitialized = false |
|||
private var model: String = "gpt-3.5-turbo" // 默认模型 |
|||
private var visionModel: String = "gpt-4-vision-preview" // 默认视觉模型 |
|||
|
|||
// OkHttp客户端用于流式请求 |
|||
private val client = OkHttpClient.Builder() |
|||
.connectTimeout(30, TimeUnit.SECONDS) |
|||
.readTimeout(30, TimeUnit.SECONDS) |
|||
.writeTimeout(30, TimeUnit.SECONDS) |
|||
.build() |
|||
|
|||
// MCP客户端 |
|||
private var mcpClient: MCPClient? = null |
|||
private var isMcpInitialized = false |
|||
|
|||
// 当前事件源 |
|||
private var currentEventSource: EventSource? = null |
|||
private var isCanceled = false |
|||
|
|||
/** |
|||
* 初始化OpenAI服务 |
|||
*/ |
|||
fun initialize(apiKey: String, baseUrl: String, model: String, mcpServer: String): Boolean { |
|||
this.apiKey = apiKey |
|||
if (baseUrl.isNotEmpty()) { |
|||
this.baseUrl = baseUrl |
|||
} |
|||
if (model.isNotEmpty()) { |
|||
this.model = model |
|||
} |
|||
|
|||
// 初始化OpenAI客户端 |
|||
try { |
|||
val timeout = Timeout(socket = 30.seconds, connect = 30.seconds, request = 30.seconds) |
|||
val host = if (baseUrl.isNotEmpty() && baseUrl != "https://api.openai.com/v1/chat/completions") { |
|||
OpenAIHost(baseUrl) |
|||
} else { |
|||
OpenAIHost.Default |
|||
} |
|||
|
|||
val config = OpenAIConfig( |
|||
token = apiKey, |
|||
host = host, |
|||
timeout = timeout |
|||
) |
|||
|
|||
openAI = OpenAI(config) |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "OpenAI客户端初始化失败: ${e.message}", e) |
|||
return false |
|||
} |
|||
|
|||
// 初始化MCPClient |
|||
initializeMcpClient(mcpServer) |
|||
|
|||
isInitialized = apiKey.isNotEmpty() && openAI != null |
|||
return isInitialized |
|||
} |
|||
|
|||
/** |
|||
* 将文件转换为Base64字符串 |
|||
*/ |
|||
fun fileToBase64(filePath: String, maxSizeKB: Int = 20480): String? { |
|||
try { |
|||
val file = File(filePath) |
|||
if (!file.exists() || !file.isFile) { |
|||
Log.e(TAG, "文件不存在: $filePath") |
|||
return null |
|||
} |
|||
|
|||
// 读取文件并压缩(如果需要) |
|||
val originalBitmap = BitmapFactory.decodeFile(filePath) |
|||
if (originalBitmap == null) { |
|||
Log.e(TAG, "无法解码图片: $filePath") |
|||
return null |
|||
} |
|||
|
|||
val outputStream = ByteArrayOutputStream() |
|||
var quality = 100 |
|||
var compressedBitmap = originalBitmap |
|||
|
|||
// 检查图片尺寸,限制最大为1024*1024 |
|||
val maxDimension = 1024 |
|||
if (originalBitmap.width > maxDimension || originalBitmap.height > maxDimension) { |
|||
Log.d(TAG, "图片尺寸超过限制,进行缩放: ${originalBitmap.width}x${originalBitmap.height} -> ${maxDimension}x${maxDimension}") |
|||
|
|||
// 计算缩放比例,保持纵横比 |
|||
val widthRatio = maxDimension.toFloat() / originalBitmap.width |
|||
val heightRatio = maxDimension.toFloat() / originalBitmap.height |
|||
val ratio = Math.min(widthRatio, heightRatio) |
|||
|
|||
val newWidth = (originalBitmap.width * ratio).toInt() |
|||
val newHeight = (originalBitmap.height * ratio).toInt() |
|||
|
|||
compressedBitmap = Bitmap.createScaledBitmap(originalBitmap, newWidth, newHeight, true) |
|||
Log.d(TAG, "缩放后图片尺寸: ${newWidth}x${newHeight}") |
|||
} |
|||
|
|||
// 如果原始图片太大,继续优化文件大小 |
|||
var fileSize = file.length() / 1024 // 转为KB |
|||
if (fileSize > maxSizeKB) { |
|||
val scale = Math.sqrt(maxSizeKB.toDouble() / fileSize) |
|||
val newWidth = (compressedBitmap.width * scale).toInt() |
|||
val newHeight = (compressedBitmap.height * scale).toInt() |
|||
compressedBitmap = Bitmap.createScaledBitmap(compressedBitmap, newWidth, newHeight, true) |
|||
quality = 85 |
|||
} |
|||
|
|||
// 压缩图片 |
|||
compressedBitmap.compress(Bitmap.CompressFormat.JPEG, quality, outputStream) |
|||
val imageBytes = outputStream.toByteArray() |
|||
|
|||
// 检查压缩后大小 |
|||
if (imageBytes.size / 1024 > maxSizeKB) { |
|||
Log.w(TAG, "压缩后图片仍然超出大小限制: ${imageBytes.size / 1024}KB > ${maxSizeKB}KB") |
|||
} |
|||
|
|||
// 转为Base64 |
|||
return Base64.encodeToString(imageBytes, Base64.NO_WRAP) |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "转换文件到Base64失败: ${e.message}", e) |
|||
return null |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 创建带图片的用户消息 |
|||
*/ |
|||
fun createUserMessageWithImage(text: String, imageBase64: String): JSONObject { |
|||
// 创建包含文本和图片的内容数组 |
|||
val contentArray = JSONArray().apply { |
|||
// 添加文本部分 |
|||
if (text.isNotEmpty()) { |
|||
put(JSONObject().apply { |
|||
put("type", "text") |
|||
put("text", text) |
|||
}) |
|||
} |
|||
|
|||
// 添加图片部分 |
|||
put(JSONObject().apply { |
|||
put("type", "image_url") |
|||
put("image_url", JSONObject().apply { |
|||
put("url", "data:image/jpeg;base64,$imageBase64") |
|||
}) |
|||
}) |
|||
} |
|||
|
|||
return JSONObject().apply { |
|||
put("role", "user") |
|||
put("content", contentArray) |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 创建用户消息 |
|||
*/ |
|||
fun createUserMessage(content: String): JSONObject { |
|||
return JSONObject().apply { |
|||
put("role", "user") |
|||
put("content", content) |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 创建系统消息 |
|||
*/ |
|||
fun createSystemMessage(content: String): JSONObject { |
|||
return JSONObject().apply { |
|||
put("role", "system") |
|||
put("content", content) |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 创建助手消息 |
|||
*/ |
|||
fun createAssistantMessage(content: String): JSONObject { |
|||
return JSONObject().apply { |
|||
put("role", "assistant") |
|||
put("content", content) |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* MCP客户端是否已初始化 |
|||
*/ |
|||
fun isMcpInitialized(): Boolean { |
|||
return isMcpInitialized && mcpClient?.isConnected() == true |
|||
} |
|||
|
|||
/** |
|||
* 关闭MCP客户端 |
|||
*/ |
|||
fun closeMcpClient() { |
|||
mcpClient?.close() |
|||
mcpClient = null |
|||
isMcpInitialized = false |
|||
} |
|||
|
|||
/** |
|||
* 初始化MCP客户端 |
|||
*/ |
|||
private fun initializeMcpClient(mcpServer: String): Boolean { |
|||
if (mcpClient != null) { |
|||
mcpClient?.close() |
|||
} |
|||
|
|||
if (mcpServer.isEmpty()) { |
|||
return false |
|||
} |
|||
|
|||
mcpClient = MCPClient(context) |
|||
|
|||
// 在后台线程中初始化MCP客户端 |
|||
launch { |
|||
try { |
|||
val result = mcpClient?.connectToSSE(mcpServer) ?: false |
|||
isMcpInitialized = result |
|||
Log.d(TAG, "MCP客户端初始化${if (result) "成功" else "失败"}") |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "MCP客户端初始化失败: ${e.message}", e) |
|||
isMcpInitialized = false |
|||
} |
|||
} |
|||
|
|||
return true // 立即返回,实际连接在后台进行 |
|||
} |
|||
|
|||
/** |
|||
* 处理MCP工具调用 |
|||
*/ |
|||
private suspend fun handleMcpToolCall(functionCall: JSONObject): JSONObject? { |
|||
if (mcpClient == null || !isMcpInitialized) { |
|||
return JSONObject().apply { put("context", "MCP客户端未初始化") } |
|||
} |
|||
|
|||
try { |
|||
// 获取函数名称 |
|||
val name = functionCall.getString("name") |
|||
|
|||
// 获取参数 |
|||
val argumentsJson = functionCall.getString("arguments") |
|||
val arguments = mcpClient?.parseJsonArguments(argumentsJson) ?: mapOf() |
|||
|
|||
// 调用工具 |
|||
return mcpClient?.callTool(name, arguments) |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "处理MCP工具调用失败: ${e.message}", e) |
|||
return JSONObject().apply { put("context", "处理MCP工具调用失败: ${e.message}") } |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 将JSONArray转换为ChatCompletionRequest中的消息列表 |
|||
*/ |
|||
private fun parseMessages(messagesArray: JSONArray): List<ChatMessage> { |
|||
val messages = mutableListOf<ChatMessage>() |
|||
|
|||
for (i in 0 until messagesArray.length()) { |
|||
val messageObj = messagesArray.getJSONObject(i) |
|||
val role = messageObj.getString("role") |
|||
|
|||
when (role) { |
|||
"system" -> { |
|||
val content = messageObj.getString("content") |
|||
messages.add(ChatMessage(role = ChatRole.System, content = content)) |
|||
} |
|||
"user" -> { |
|||
// 检查是否有多媒体内容 |
|||
if (messageObj.has("content") && messageObj.get("content") is JSONArray) { |
|||
val contentArray = messageObj.getJSONArray("content") |
|||
val parts = mutableListOf<ChatMessageContent>() |
|||
|
|||
for (j in 0 until contentArray.length()) { |
|||
val contentObj = contentArray.getJSONObject(j) |
|||
val type = contentObj.getString("type") |
|||
|
|||
when (type) { |
|||
"text" -> { |
|||
parts.add(TextContent(contentObj.getString("text"))) |
|||
} |
|||
"image_url" -> { |
|||
val imageUrlObj = contentObj.getJSONObject("image_url") |
|||
val url = imageUrlObj.getString("url") |
|||
parts.add(ImageContent(url)) |
|||
} |
|||
} |
|||
} |
|||
|
|||
messages.add(ChatMessage( |
|||
role = ChatRole.User, |
|||
content = parts |
|||
)) |
|||
} else { |
|||
// 普通文本消息 |
|||
val content = messageObj.getString("content") |
|||
messages.add(ChatMessage(role = ChatRole.User, content = content)) |
|||
} |
|||
} |
|||
"assistant" -> { |
|||
if (messageObj.has("tool_calls")) { |
|||
// 处理工具调用 |
|||
val toolCalls = messageObj.getJSONArray("tool_calls") |
|||
val toolCallsList = mutableListOf<ToolCall>() |
|||
|
|||
for (j in 0 until toolCalls.length()) { |
|||
val toolCall = toolCalls.getJSONObject(j) |
|||
val id = toolCall.getString("id") |
|||
val function = toolCall.getJSONObject("function") |
|||
val name = function.getString("name") |
|||
val arguments = function.getString("arguments") |
|||
|
|||
toolCallsList.add(ToolCall( |
|||
id = id, |
|||
type = ToolCallType.Function, |
|||
function = FunctionCall( |
|||
name = name, |
|||
arguments = arguments |
|||
) |
|||
)) |
|||
} |
|||
|
|||
val content = if (messageObj.has("content")) messageObj.getString("content") else "" |
|||
|
|||
messages.add(ChatMessage( |
|||
role = ChatRole.Assistant, |
|||
content = content, |
|||
toolCalls = toolCallsList |
|||
)) |
|||
} else { |
|||
// 普通消息 |
|||
val content = messageObj.getString("content") |
|||
messages.add(ChatMessage(role = ChatRole.Assistant, content = content)) |
|||
} |
|||
} |
|||
"tool" -> { |
|||
val content = messageObj.getString("content") |
|||
val toolCallId = messageObj.getString("tool_call_id") |
|||
messages.add(ChatMessage( |
|||
role = ChatRole.Tool, |
|||
content = content, |
|||
toolCallId = toolCallId |
|||
)) |
|||
} |
|||
} |
|||
} |
|||
|
|||
return messages |
|||
} |
|||
|
|||
/** |
|||
* 从MCP客户端获取工具定义 |
|||
*/ |
|||
private fun getToolsFromMcpClient(): List<Tool> { |
|||
val tools = mutableListOf<Tool>() |
|||
|
|||
mcpClient?.getToolMaps()?.forEach { toolMap -> |
|||
try { |
|||
val toolJson = JSONObject(toolMap) |
|||
val name = toolJson.getString("name") |
|||
val description = toolJson.getString("description") |
|||
|
|||
val parametersJson = if (toolJson.has("parameters")) toolJson.getJSONObject("parameters") else null |
|||
val parameterProperties = mutableMapOf<String, ParameterDefinition>() |
|||
val requiredParams = mutableListOf<String>() |
|||
|
|||
if (parametersJson != null && parametersJson.has("properties")) { |
|||
val properties = parametersJson.getJSONObject("properties") |
|||
val keys = properties.keys() |
|||
|
|||
while (keys.hasNext()) { |
|||
val key = keys.next() |
|||
val propertyObj = properties.getJSONObject(key) |
|||
val propType = propertyObj.optString("type", "string") |
|||
val propDescription = propertyObj.optString("description", "") |
|||
|
|||
parameterProperties[key] = ParameterDefinition( |
|||
type = propType, |
|||
description = propDescription |
|||
) |
|||
} |
|||
|
|||
// 获取必填参数 |
|||
if (parametersJson.has("required")) { |
|||
val requiredArr = parametersJson.getJSONArray("required") |
|||
for (i in 0 until requiredArr.length()) { |
|||
requiredParams.add(requiredArr.getString(i)) |
|||
} |
|||
} |
|||
} |
|||
|
|||
tools.add(Tool( |
|||
type = ToolType.Function, |
|||
function = FunctionDefinition( |
|||
name = name, |
|||
description = description, |
|||
parameters = FunctionParameters( |
|||
type = "object", |
|||
properties = parameterProperties, |
|||
required = requiredParams |
|||
) |
|||
) |
|||
)) |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "解析工具定义失败: ${e.message}", e) |
|||
} |
|||
} |
|||
|
|||
return tools |
|||
} |
|||
|
|||
/** |
|||
* 发送消息(非流式输出) |
|||
*/ |
|||
suspend fun sendMessage(messages: JSONArray): String { |
|||
if (!isInitialized || apiKey.isEmpty() || openAI == null) { |
|||
throw IOException("OpenAI服务未初始化") |
|||
} |
|||
|
|||
try { |
|||
// 解析消息 |
|||
val parsedMessages = parseMessages(messages) |
|||
|
|||
// 获取工具列表 |
|||
val tools = getToolsFromMcpClient() |
|||
|
|||
// 创建请求 |
|||
val request = ChatCompletionRequest( |
|||
model = ModelId(model), |
|||
messages = parsedMessages, |
|||
tools = if (tools.isNotEmpty()) tools else null, |
|||
temperature = 0.7, |
|||
maxTokens = 2000 |
|||
) |
|||
|
|||
// 发送请求 |
|||
val response = openAI!!.chatCompletion(request) |
|||
|
|||
// 解析响应 |
|||
val choice = response.choices.firstOrNull() ?: throw IOException("无效的响应格式") |
|||
|
|||
// 检查是否有工具调用 |
|||
if (choice.message.toolCalls?.isNotEmpty() == true) { |
|||
val toolCall = choice.message.toolCalls?.first() |
|||
if (toolCall != null && toolCall.function != null) { |
|||
val functionCall = JSONObject().apply { |
|||
put("name", toolCall.function.name) |
|||
put("arguments", toolCall.function.arguments) |
|||
put("id", toolCall.id) |
|||
} |
|||
return functionCall.toString() |
|||
} |
|||
} |
|||
|
|||
// 返回消息内容 |
|||
return choice.message.content ?: "" |
|||
|
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "发送消息失败: ${e.message}", e) |
|||
throw IOException("与AI服务通信失败: ${e.message}") |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 发送消息(流式输出) |
|||
*/ |
|||
fun sendMessageStream(messages: JSONArray, callback: StreamCallback) { |
|||
if (!isInitialized || apiKey.isEmpty()) { |
|||
// 在主线程执行回调 |
|||
launch { |
|||
withContext(Dispatchers.Main) { |
|||
callback.onError(Exception("OpenAI服务未初始化")) |
|||
} |
|||
} |
|||
return |
|||
} |
|||
|
|||
// 重置取消状态 |
|||
isCanceled = false |
|||
|
|||
// 检查是否有图片消息 |
|||
var hasImageContent = false |
|||
var currentModel = model |
|||
|
|||
for (i in 0 until messages.length()) { |
|||
val messageObj = messages.getJSONObject(i) |
|||
if (messageObj.getString("role") == "user" && messageObj.has("content")) { |
|||
val content = messageObj.get("content") |
|||
if (content is JSONArray) { |
|||
for (j in 0 until content.length()) { |
|||
val contentObj = content.getJSONObject(j) |
|||
if (contentObj.getString("type") == "image_url") { |
|||
hasImageContent = true |
|||
currentModel = visionModel |
|||
break |
|||
} |
|||
} |
|||
} |
|||
} |
|||
if (hasImageContent) break |
|||
} |
|||
|
|||
// 构建JSON请求体 |
|||
val requestBody = JSONObject().apply { |
|||
put("model", currentModel) |
|||
put("messages", messages) |
|||
put("temperature", 0.7) |
|||
put("max_tokens", 2000) |
|||
put("stream", true) |
|||
|
|||
// 添加工具列表 |
|||
val tools = JSONArray() |
|||
|
|||
// 使用MCPClient提供的所有工具 |
|||
mcpClient?.getToolMaps()?.forEach { toolMap -> |
|||
try { |
|||
val tool = JSONObject(toolMap) |
|||
tools.put(tool) |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "转换工具失败: ${e.message}", e) |
|||
} |
|||
} |
|||
|
|||
// 如果有工具,则添加到请求中 |
|||
if (tools.length() > 0) { |
|||
put("tools", tools) |
|||
} |
|||
} |
|||
|
|||
val mediaType = "application/json".toMediaTypeOrNull() |
|||
val request = Request.Builder() |
|||
.url(baseUrl) |
|||
.addHeader("Content-Type", "application/json") |
|||
.addHeader("Authorization", "Bearer $apiKey") |
|||
.addHeader("Accept", "text/event-stream") |
|||
.post(requestBody.toString().toRequestBody(mediaType)) |
|||
.build() |
|||
|
|||
// 创建事件源 |
|||
val factory = EventSources.createFactory(client) |
|||
|
|||
// 工具调用相关变量 |
|||
val toolCalls = mutableMapOf<Int, ToolCallInfo>() |
|||
|
|||
val eventSourceListener = object : EventSourceListener() { |
|||
override fun onOpen(eventSource: EventSource, response: Response) { |
|||
Log.d(TAG, "SSE连接已打开") |
|||
} |
|||
|
|||
override fun onEvent(eventSource: EventSource, id: String?, type: String?, data: String) { |
|||
if (isCanceled) return |
|||
|
|||
if (data == "[DONE]" || data == "[\"DONE\"]") { |
|||
// 处理可能的工具调用 |
|||
processToolCalls(toolCalls, callback, messages) |
|||
return |
|||
} |
|||
|
|||
try { |
|||
val jsonData = JSONObject(data) |
|||
|
|||
// 处理消息内容 |
|||
if (jsonData.has("choices")) { |
|||
val choices = jsonData.getJSONArray("choices") |
|||
if (choices.length() > 0) { |
|||
val choice = choices.getJSONObject(0) |
|||
|
|||
// 处理delta内容 |
|||
if (choice.has("delta")) { |
|||
val delta = choice.getJSONObject("delta") |
|||
|
|||
// 处理普通文本内容 |
|||
if (delta.has("content")) { |
|||
val content = delta.getString("content") |
|||
if (!isCanceled) { |
|||
launch { |
|||
withContext(Dispatchers.Main) { |
|||
callback.onToken(content) |
|||
} |
|||
} |
|||
} |
|||
} |
|||
|
|||
// 处理工具调用 |
|||
if (delta.has("tool_calls")) { |
|||
val deltaToolCalls = delta.getJSONArray("tool_calls") |
|||
for (i in 0 until deltaToolCalls.length()) { |
|||
val toolCall = deltaToolCalls.getJSONObject(i) |
|||
val index = toolCall.getInt("index") |
|||
|
|||
// 创建或获取现有的工具调用信息 |
|||
val toolCallInfo = toolCalls.getOrPut(index) { ToolCallInfo() } |
|||
|
|||
// 更新ID |
|||
if (toolCall.has("id")) { |
|||
toolCallInfo.id = toolCall.getString("id") |
|||
} |
|||
|
|||
// 更新函数信息 |
|||
if (toolCall.has("function")) { |
|||
val function = toolCall.getJSONObject("function") |
|||
|
|||
if (function.has("name")) { |
|||
toolCallInfo.name = function.getString("name") |
|||
} |
|||
|
|||
if (function.has("arguments")) { |
|||
toolCallInfo.arguments += function.getString("arguments") |
|||
} |
|||
} |
|||
} |
|||
} |
|||
} |
|||
|
|||
// 检查是否有表示完成的标志 |
|||
if (choice.has("finish_reason")) { |
|||
val finishReason = choice.getString("finish_reason") |
|||
if (finishReason == "stop" || finishReason == "length") { |
|||
// 正常完成,没有工具调用 |
|||
launch { |
|||
withContext(Dispatchers.Main) { |
|||
callback.onComplete() |
|||
} |
|||
} |
|||
closeEventSource() |
|||
} else if (finishReason == "tool_calls") { |
|||
// 处理工具调用 |
|||
processToolCalls(toolCalls, callback, messages) |
|||
closeEventSource() |
|||
} |
|||
} |
|||
} |
|||
} |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "解析事件数据失败: ${e.message}", e) |
|||
} |
|||
} |
|||
|
|||
override fun onClosed(eventSource: EventSource) { |
|||
Log.d(TAG, "SSE连接已关闭") |
|||
if (!isCanceled) { |
|||
// 如果没有正常完成,但连接关闭了,则处理最后可能的工具调用 |
|||
if (toolCalls.isNotEmpty()) { |
|||
processToolCalls(toolCalls, callback, messages) |
|||
} else { |
|||
launch { |
|||
withContext(Dispatchers.Main) { |
|||
callback.onComplete() |
|||
} |
|||
} |
|||
} |
|||
} |
|||
} |
|||
|
|||
override fun onFailure(eventSource: EventSource, t: Throwable?, response: Response?) { |
|||
if (isCanceled) return |
|||
|
|||
val errorCode = response?.code ?: 0 |
|||
val errorMessage = t?.message ?: "未知错误" |
|||
Log.e(TAG, "SSE连接失败: $errorCode - $errorMessage") |
|||
|
|||
launch { |
|||
withContext(Dispatchers.Main) { |
|||
callback.onError(Exception("流式请求失败: $errorMessage")) |
|||
} |
|||
} |
|||
closeEventSource() |
|||
} |
|||
} |
|||
|
|||
currentEventSource = factory.newEventSource(request, eventSourceListener) |
|||
} |
|||
|
|||
/** |
|||
* 处理工具调用 |
|||
*/ |
|||
private fun processToolCalls(toolCalls: Map<Int, ToolCallInfo>, callback: StreamCallback, messages: JSONArray? = null): Boolean { |
|||
if (toolCalls.isEmpty()) return false |
|||
|
|||
val firstToolCall = toolCalls.entries.firstOrNull()?.value ?: return false |
|||
|
|||
if (firstToolCall.isValid()) { |
|||
// 创建函数调用JSON对象 |
|||
val functionCall = JSONObject().apply { |
|||
put("name", firstToolCall.name) |
|||
put("arguments", firstToolCall.arguments) |
|||
put("id", firstToolCall.id) |
|||
} |
|||
Log.d(TAG, "工具调用: $functionCall") |
|||
|
|||
// 通知上层回调 |
|||
launch { |
|||
withContext(Dispatchers.Main) { |
|||
callback.onFunctionCall(functionCall) |
|||
} |
|||
} |
|||
|
|||
// 在协程中自动处理工具调用 |
|||
if (messages != null) { |
|||
launch { |
|||
try { |
|||
if (!isCanceled) { |
|||
handleToolCall(functionCall, messages, callback) |
|||
} |
|||
} catch (e: Exception) { |
|||
if (!isCanceled) { |
|||
Log.e(TAG, "处理工具调用时发生异常: ${e.message}", e) |
|||
try { |
|||
val errorMessage = "工具调用处理失败: ${e.message}" |
|||
sendFunctionCallResult( |
|||
messages = messages, |
|||
functionCall = functionCall, |
|||
functionResult = errorMessage, |
|||
callback = callback |
|||
) |
|||
} catch (e2: Exception) { |
|||
Log.e(TAG, "发送工具调用错误结果失败: ${e2.message}", e2) |
|||
withContext(Dispatchers.Main) { |
|||
callback.onError(Exception("工具调用处理失败: ${e.message}")) |
|||
} |
|||
} |
|||
} |
|||
} |
|||
} |
|||
} |
|||
return true |
|||
} |
|||
return false |
|||
} |
|||
|
|||
/** |
|||
* 处理工具调用 |
|||
*/ |
|||
private suspend fun handleToolCall(functionCall: JSONObject, messages: JSONArray, callback: StreamCallback) { |
|||
try { |
|||
// 获取函数名称和参数 |
|||
val name = functionCall.getString("name") |
|||
val argumentsJson = functionCall.getString("arguments") |
|||
|
|||
Log.d(TAG, "处理工具调用: name=$name, arguments=$argumentsJson") |
|||
|
|||
// 调用MCP工具 |
|||
val result = handleMcpToolCall(functionCall) |
|||
|
|||
// 检查是否已取消 |
|||
if (isCanceled) { |
|||
Log.d(TAG, "工具调用已被取消,不处理结果") |
|||
return |
|||
} |
|||
|
|||
// 回调结果 |
|||
val resultObj = result ?: JSONObject().apply { put("context", "工具调用失败") } |
|||
withContext(Dispatchers.Main) { |
|||
callback.onFunctionCallResult(functionCall, resultObj) |
|||
} |
|||
|
|||
// 从结果中提取内容 |
|||
val resultContent = if (resultObj.has("context") && resultObj.optString("context").isNotEmpty()) { |
|||
resultObj.getString("context") |
|||
} else { |
|||
val jsonString = resultObj.toString() |
|||
if (jsonString == "{}") "工具调用失败" else jsonString |
|||
} |
|||
|
|||
// 发送函数调用结果 |
|||
sendFunctionCallResult(messages, functionCall, resultContent, callback) |
|||
} catch (e: Exception) { |
|||
if (!isCanceled) { |
|||
Log.e(TAG, "处理工具调用失败: ${e.message}", e) |
|||
withContext(Dispatchers.Main) { |
|||
callback.onError(Exception("工具调用失败: ${e.message}")) |
|||
} |
|||
} |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 发送函数调用结果 |
|||
*/ |
|||
fun sendFunctionCallResult( |
|||
messages: JSONArray, |
|||
functionCall: JSONObject, |
|||
functionResult: String, |
|||
callback: StreamCallback |
|||
) { |
|||
if (isCanceled) { |
|||
Log.d(TAG, "请求已取消,不发送函数调用结果") |
|||
return |
|||
} |
|||
|
|||
try { |
|||
val fullMessages = JSONArray() |
|||
|
|||
// 添加原始消息 |
|||
for (i in 0 until messages.length()) { |
|||
fullMessages.put(messages.getJSONObject(i)) |
|||
} |
|||
|
|||
// 添加函数调用消息 |
|||
val callId = functionCall.optString("id", "call_${System.currentTimeMillis()}") |
|||
fullMessages.put(JSONObject().apply { |
|||
put("role", "assistant") |
|||
put("content", "") |
|||
|
|||
// 添加工具调用 |
|||
val toolCalls = JSONArray().apply { |
|||
val toolCall = JSONObject().apply { |
|||
put("id", callId) |
|||
put("type", "function") |
|||
put("function", JSONObject().apply { |
|||
put("name", functionCall.getString("name")) |
|||
put("arguments", functionCall.getString("arguments")) |
|||
}) |
|||
} |
|||
put(toolCall) |
|||
} |
|||
put("tool_calls", toolCalls) |
|||
}) |
|||
|
|||
// 添加函数调用结果 |
|||
fullMessages.put(JSONObject().apply { |
|||
put("role", "tool") |
|||
put("content", functionResult) |
|||
put("tool_call_id", callId) |
|||
}) |
|||
|
|||
// 发送完整对话 |
|||
sendMessageStream(fullMessages, callback) |
|||
|
|||
} catch (e: Exception) { |
|||
if (!isCanceled) { |
|||
Log.e(TAG, "发送函数调用结果失败: ${e.message}", e) |
|||
launch { |
|||
withContext(Dispatchers.Main) { |
|||
callback.onError(Exception("发送函数调用结果失败: ${e.message}")) |
|||
} |
|||
} |
|||
} |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 关闭事件源 |
|||
*/ |
|||
private fun closeEventSource() { |
|||
currentEventSource?.let { |
|||
try { |
|||
it.cancel() |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "关闭事件源失败: ${e.message}", e) |
|||
} |
|||
currentEventSource = null |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 取消当前流式请求 |
|||
*/ |
|||
fun cancelCurrentStream(): Boolean { |
|||
isCanceled = true |
|||
closeEventSource() |
|||
return true |
|||
} |
|||
|
|||
/** |
|||
* 释放资源 |
|||
*/ |
|||
fun dispose() { |
|||
job.cancel() // 取消所有协程 |
|||
cancelCurrentStream() // 取消当前流式请求 |
|||
closeMcpClient() // 关闭MCP客户端 |
|||
} |
|||
|
|||
/** |
|||
* 工具调用信息类 |
|||
*/ |
|||
private class ToolCallInfo { |
|||
var id: String = "" |
|||
var name: String = "" |
|||
var arguments: String = "" |
|||
|
|||
fun isValid(): Boolean { |
|||
return id.isNotEmpty() && name.isNotEmpty() |
|||
} |
|||
} |
|||
|
|||
/** |
|||
* 流式回调接口 |
|||
*/ |
|||
interface StreamCallback { |
|||
fun onToken(token: String) |
|||
fun onComplete() |
|||
fun onError(e: Exception) |
|||
fun onFunctionCall(functionCall: JSONObject) |
|||
fun onFunctionCallResult(functionCall: JSONObject, functionCallResult: JSONObject) |
|||
} |
|||
|
|||
/** |
|||
* 参数定义 |
|||
*/ |
|||
private data class ParameterDefinition( |
|||
val type: String, |
|||
val description: String |
|||
) |
|||
|
|||
/** |
|||
* 注册函数 |
|||
*/ |
|||
fun registerFunction(name: String, description: String, parameters: JSONObject): Boolean { |
|||
try { |
|||
// 确保MCPClient已初始化 |
|||
if (mcpClient == null) { |
|||
mcpClient = MCPClient(context) |
|||
} |
|||
|
|||
// 创建函数处理器 |
|||
val handler = object : FunctionHandler { |
|||
override suspend fun handle(arguments: Map<String, Any>): String { |
|||
// 由于本地函数的实际处理是在Flutter端完成的 |
|||
// 这里只需返回一个标记,表示该函数是本地函数 |
|||
return "LOCAL_FUNCTION:$name" |
|||
} |
|||
} |
|||
|
|||
// 注册本地函数 |
|||
return mcpClient?.registerLocalFunction(name, description, parameters, handler) ?: false |
|||
} catch (e: Exception) { |
|||
Log.e(TAG, "注册函数失败: ${e.message}", e) |
|||
return false |
|||
} |
|||
} |
|||
} |
|||
@ -0,0 +1,149 @@ |
|||
package com.yunqiinnovation.open_ai |
|||
|
|||
import android.content.Context |
|||
import androidx.annotation.NonNull |
|||
import io.flutter.embedding.engine.plugins.FlutterPlugin |
|||
import io.flutter.plugin.common.MethodCall |
|||
import io.flutter.plugin.common.MethodChannel |
|||
import io.flutter.plugin.common.MethodChannel.MethodCallHandler |
|||
import io.flutter.plugin.common.MethodChannel.Result |
|||
import org.json.JSONArray |
|||
import org.json.JSONObject |
|||
import kotlinx.coroutines.* |
|||
|
|||
class OpenAiPlugin : FlutterPlugin, MethodCallHandler { |
|||
private lateinit var channel: MethodChannel |
|||
private lateinit var context: Context |
|||
private lateinit var openAIService: OpenAIService |
|||
private val scope = CoroutineScope(Dispatchers.IO + SupervisorJob()) |
|||
|
|||
override fun onAttachedToEngine(@NonNull flutterPluginBinding: FlutterPlugin.FlutterPluginBinding) { |
|||
channel = MethodChannel(flutterPluginBinding.binaryMessenger, "com.yunqiinnovation.open_ai") |
|||
context = flutterPluginBinding.applicationContext |
|||
openAIService = OpenAIService(context) |
|||
channel.setMethodCallHandler(this) |
|||
} |
|||
|
|||
override fun onMethodCall(@NonNull call: MethodCall, @NonNull result: Result) { |
|||
when (call.method) { |
|||
"initialize" -> { |
|||
val apiKey = call.argument<String>("apiKey") ?: "" |
|||
val baseUrl = call.argument<String>("baseUrl") ?: "" |
|||
val model = call.argument<String>("model") ?: "" |
|||
val mcpServer = call.argument<String>("mcpServer") ?: "" |
|||
val success = openAIService.initialize(apiKey, baseUrl, model, mcpServer) |
|||
result.success(success) |
|||
} |
|||
"createUserMessage" -> { |
|||
val content = call.argument<String>("content") ?: "" |
|||
val message = openAIService.createUserMessage(content) |
|||
result.success(message.toString()) |
|||
} |
|||
"createSystemMessage" -> { |
|||
val content = call.argument<String>("content") ?: "" |
|||
val message = openAIService.createSystemMessage(content) |
|||
result.success(message.toString()) |
|||
} |
|||
"createAssistantMessage" -> { |
|||
val content = call.argument<String>("content") ?: "" |
|||
val message = openAIService.createAssistantMessage(content) |
|||
result.success(message.toString()) |
|||
} |
|||
"createUserMessageWithImage" -> { |
|||
val text = call.argument<String>("text") ?: "" |
|||
val imageBase64 = call.argument<String>("imageBase64") ?: "" |
|||
val message = openAIService.createUserMessageWithImage(text, imageBase64) |
|||
result.success(message.toString()) |
|||
} |
|||
"sendMessage" -> { |
|||
val messagesJson = call.argument<String>("messages") ?: "[]" |
|||
val messages = JSONArray(messagesJson) |
|||
|
|||
scope.launch { |
|||
try { |
|||
val response = openAIService.sendMessage(messages) |
|||
withContext(Dispatchers.Main) { |
|||
result.success(response) |
|||
} |
|||
} catch (e: Exception) { |
|||
withContext(Dispatchers.Main) { |
|||
result.error("OPENAI_ERROR", e.message ?: "Unknown error", null) |
|||
} |
|||
} |
|||
} |
|||
} |
|||
"sendMessageStream" -> { |
|||
val messagesJson = call.argument<String>("messages") ?: "[]" |
|||
val messages = JSONArray(messagesJson) |
|||
val streamId = call.argument<String>("streamId") ?: "${System.currentTimeMillis()}" |
|||
|
|||
val callback = object : OpenAIService.StreamCallback { |
|||
override fun onToken(token: String) { |
|||
val map = mapOf( |
|||
"type" to "token", |
|||
"streamId" to streamId, |
|||
"data" to token |
|||
) |
|||
channel.invokeMethod("onStreamEvent", map) |
|||
} |
|||
|
|||
override fun onComplete() { |
|||
val map = mapOf( |
|||
"type" to "complete", |
|||
"streamId" to streamId |
|||
) |
|||
channel.invokeMethod("onStreamEvent", map) |
|||
} |
|||
|
|||
override fun onError(e: Exception) { |
|||
val map = mapOf( |
|||
"type" to "error", |
|||
"streamId" to streamId, |
|||
"error" to (e.message ?: "Unknown error") |
|||
) |
|||
channel.invokeMethod("onStreamEvent", map) |
|||
} |
|||
|
|||
override fun onFunctionCall(functionCall: JSONObject) { |
|||
val map = mapOf( |
|||
"type" to "functionCall", |
|||
"streamId" to streamId, |
|||
"data" to functionCall.toString() |
|||
) |
|||
channel.invokeMethod("onStreamEvent", map) |
|||
} |
|||
|
|||
override fun onFunctionCallResult(functionCall: JSONObject, functionCallResult: JSONObject) { |
|||
val map = mapOf( |
|||
"type" to "functionCallResult", |
|||
"streamId" to streamId, |
|||
"functionCall" to functionCall.toString(), |
|||
"result" to functionCallResult.toString() |
|||
) |
|||
channel.invokeMethod("onStreamEvent", map) |
|||
} |
|||
} |
|||
|
|||
openAIService.sendMessageStream(messages, callback) |
|||
result.success(streamId) |
|||
} |
|||
"cancelCurrentStream" -> { |
|||
val success = openAIService.cancelCurrentStream() |
|||
result.success(success) |
|||
} |
|||
"dispose" -> { |
|||
openAIService.dispose() |
|||
result.success(null) |
|||
} |
|||
else -> { |
|||
result.notImplemented() |
|||
} |
|||
} |
|||
} |
|||
|
|||
override fun onDetachedFromEngine(@NonNull binding: FlutterPlugin.FlutterPluginBinding) { |
|||
channel.setMethodCallHandler(null) |
|||
scope.cancel() // Cancel all coroutines when the plugin is detached |
|||
openAIService.dispose() // Clean up resources |
|||
} |
|||
} |
|||
@ -0,0 +1,184 @@ |
|||
import 'dart:async'; |
|||
import 'dart:convert'; |
|||
|
|||
import 'package:flutter/services.dart'; |
|||
|
|||
/// OpenAI和MCP集成插件 |
|||
class OpenAI { |
|||
/// 插件通道 |
|||
static const MethodChannel _channel = |
|||
MethodChannel('com.yunqiinnovation.open_ai'); |
|||
|
|||
/// 流式输出事件回调 |
|||
static final Map<String, StreamCallback> _streamCallbacks = {}; |
|||
|
|||
/// 构造函数 |
|||
OpenAI() { |
|||
_channel.setMethodCallHandler(_handleMethodCall); |
|||
} |
|||
|
|||
/// 处理来自原生端的方法调用 |
|||
Future<dynamic> _handleMethodCall(MethodCall call) async { |
|||
if (call.method == 'onStreamEvent') { |
|||
final Map<String, dynamic> args = Map<String, dynamic>.from(call.arguments); |
|||
final String streamId = args['streamId']; |
|||
final String type = args['type']; |
|||
|
|||
final callback = _streamCallbacks[streamId]; |
|||
if (callback != null) { |
|||
switch (type) { |
|||
case 'token': |
|||
callback.onToken(args['data']); |
|||
break; |
|||
case 'complete': |
|||
callback.onComplete(); |
|||
_streamCallbacks.remove(streamId); |
|||
break; |
|||
case 'error': |
|||
callback.onError(Exception(args['error'])); |
|||
_streamCallbacks.remove(streamId); |
|||
break; |
|||
case 'functionCall': |
|||
final functionCall = jsonDecode(args['data']); |
|||
callback.onFunctionCall(functionCall); |
|||
break; |
|||
case 'functionCallResult': |
|||
final functionCall = jsonDecode(args['functionCall']); |
|||
final functionCallResult = jsonDecode(args['result']); |
|||
callback.onFunctionCallResult(functionCall, functionCallResult); |
|||
break; |
|||
} |
|||
} |
|||
} |
|||
return null; |
|||
} |
|||
|
|||
/// 初始化OpenAI服务 |
|||
/// |
|||
/// [apiKey] OpenAI API密钥 |
|||
/// [baseUrl] API基础URL,可选 |
|||
/// [model] 模型名称,可选 |
|||
/// [mcpServer] MCP服务器URL,可选 |
|||
Future<bool> initialize({ |
|||
required String apiKey, |
|||
String baseUrl = '', |
|||
String model = '', |
|||
String mcpServer = '', |
|||
}) async { |
|||
final result = await _channel.invokeMethod<bool>('initialize', { |
|||
'apiKey': apiKey, |
|||
'baseUrl': baseUrl, |
|||
'model': model, |
|||
'mcpServer': mcpServer, |
|||
}); |
|||
return result ?? false; |
|||
} |
|||
|
|||
/// 创建用户消息 |
|||
Future<Map<String, dynamic>> createUserMessage(String content) async { |
|||
final result = await _channel.invokeMethod<String>('createUserMessage', { |
|||
'content': content, |
|||
}); |
|||
return jsonDecode(result ?? '{}'); |
|||
} |
|||
|
|||
/// 创建系统消息 |
|||
Future<Map<String, dynamic>> createSystemMessage(String content) async { |
|||
final result = await _channel.invokeMethod<String>('createSystemMessage', { |
|||
'content': content, |
|||
}); |
|||
return jsonDecode(result ?? '{}'); |
|||
} |
|||
|
|||
/// 创建助手消息 |
|||
Future<Map<String, dynamic>> createAssistantMessage(String content) async { |
|||
final result = await _channel.invokeMethod<String>('createAssistantMessage', { |
|||
'content': content, |
|||
}); |
|||
return jsonDecode(result ?? '{}'); |
|||
} |
|||
|
|||
/// 创建带图片的用户消息 |
|||
Future<Map<String, dynamic>> createUserMessageWithImage( |
|||
String text, |
|||
String imageBase64, |
|||
) async { |
|||
final result = await _channel.invokeMethod<String>( |
|||
'createUserMessageWithImage', |
|||
{ |
|||
'text': text, |
|||
'imageBase64': imageBase64, |
|||
}, |
|||
); |
|||
return jsonDecode(result ?? '{}'); |
|||
} |
|||
|
|||
/// 发送消息(非流式输出) |
|||
Future<String> sendMessage(List<Map<String, dynamic>> messages) async { |
|||
final messagesJson = jsonEncode(messages); |
|||
return await _channel.invokeMethod('sendMessage', { |
|||
'messages': messagesJson, |
|||
}); |
|||
} |
|||
|
|||
/// 发送消息(流式输出) |
|||
Future<String> sendMessageStream( |
|||
List<Map<String, dynamic>> messages, |
|||
StreamCallback callback, |
|||
) async { |
|||
final messagesJson = jsonEncode(messages); |
|||
final streamId = DateTime.now().millisecondsSinceEpoch.toString(); |
|||
|
|||
// 注册回调 |
|||
_streamCallbacks[streamId] = callback; |
|||
|
|||
final result = await _channel.invokeMethod<String>('sendMessageStream', { |
|||
'messages': messagesJson, |
|||
'streamId': streamId, |
|||
}); |
|||
|
|||
return result ?? streamId; |
|||
} |
|||
|
|||
/// 取消当前流式请求 |
|||
Future<bool> cancelCurrentStream() async { |
|||
final result = await _channel.invokeMethod<bool>('cancelCurrentStream'); |
|||
return result ?? false; |
|||
} |
|||
|
|||
/// 释放资源 |
|||
Future<void> dispose() async { |
|||
await _channel.invokeMethod('dispose'); |
|||
_streamCallbacks.clear(); |
|||
} |
|||
} |
|||
|
|||
/// 流式输出回调接口 |
|||
class StreamCallback { |
|||
/// 收到令牌 |
|||
final void Function(String token) onToken; |
|||
|
|||
/// 完成回调 |
|||
final void Function() onComplete; |
|||
|
|||
/// 错误回调 |
|||
final void Function(Exception e) onError; |
|||
|
|||
/// 函数调用回调 |
|||
final void Function(Map<String, dynamic> functionCall) onFunctionCall; |
|||
|
|||
/// 函数调用结果回调 |
|||
final void Function( |
|||
Map<String, dynamic> functionCall, |
|||
Map<String, dynamic> functionCallResult, |
|||
) onFunctionCallResult; |
|||
|
|||
/// 构造函数 |
|||
StreamCallback({ |
|||
required this.onToken, |
|||
required this.onComplete, |
|||
required this.onError, |
|||
required this.onFunctionCall, |
|||
required this.onFunctionCallResult, |
|||
}); |
|||
} |
|||
@ -0,0 +1,27 @@ |
|||
name: open_ai |
|||
description: OpenAI API与MCP集成插件,提供对OpenAI API的访问和MCP工具调用功能 |
|||
version: 0.1.0 |
|||
homepage: https://github.com/yunqiinnovation/deep_voice |
|||
|
|||
environment: |
|||
sdk: '>=2.18.0 <4.0.0' |
|||
flutter: ">=3.3.0" |
|||
|
|||
dependencies: |
|||
flutter: |
|||
sdk: flutter |
|||
|
|||
dev_dependencies: |
|||
flutter_test: |
|||
sdk: flutter |
|||
flutter_lints: ^2.0.0 |
|||
|
|||
# 插件平台配置 |
|||
flutter: |
|||
plugin: |
|||
platforms: |
|||
android: |
|||
package: com.yunqiinnovation.open_ai |
|||
pluginClass: OpenAiPlugin |
|||
ios: |
|||
pluginClass: OpenAiPlugin |
|||
Binary file not shown.
Binary file not shown.
@ -1,2 +1,2 @@ |
|||
#Sat May 10 17:59:51 IST 2025 |
|||
#Sun May 11 13:39:53 IST 2025 |
|||
gradle.version=8.10 |
|||
|
|||
Loading…
Reference in new issue