20 changed files with 1460 additions and 324 deletions
@ -0,0 +1,12 @@ |
|||||
|
{ |
||||
|
"mcpServers": { |
||||
|
|
||||
|
"amap-amap-sse": { |
||||
|
"url": "https://mcp.amap.com/sse?key=e5fdc9605eabdeb5626f18f5721f343d" |
||||
|
}, |
||||
|
"web-search": { |
||||
|
"url": "http://mcp.ideapsound.com:8000/sse" |
||||
|
} |
||||
|
|
||||
|
} |
||||
|
} |
||||
@ -0,0 +1,357 @@ |
|||||
|
package com.yunqiinnovation.open_ai_service.mcp |
||||
|
|
||||
|
import android.util.Log |
||||
|
import io.ktor.client.* |
||||
|
import io.ktor.client.plugins.sse.* |
||||
|
import io.ktor.client.request.* |
||||
|
import io.ktor.client.statement.* |
||||
|
import io.ktor.http.* |
||||
|
import io.modelcontextprotocol.kotlin.sdk.JSONRPCMessage |
||||
|
import io.modelcontextprotocol.kotlin.sdk.shared.AbstractTransport |
||||
|
import kotlinx.coroutines.* |
||||
|
import kotlinx.serialization.encodeToString |
||||
|
import kotlinx.serialization.json.Json |
||||
|
import kotlinx.serialization.decodeFromString |
||||
|
import kotlin.properties.Delegates |
||||
|
import kotlin.time.Duration |
||||
|
import java.util.concurrent.atomic.AtomicBoolean |
||||
|
import org.json.JSONObject |
||||
|
|
||||
|
/** |
||||
|
* 自定义SSE客户端传输层,修复原始SseClientTransport中的URL拼接问题 |
||||
|
* 解决URL查询参数与路径拼接错误的问题,确保消息端点URL格式正确 |
||||
|
*/ |
||||
|
class CustomSseClientTransport( |
||||
|
private val client: HttpClient, |
||||
|
private val urlString: String?, |
||||
|
private val reconnectionTime: Duration? = null, |
||||
|
private val requestBuilder: HttpRequestBuilder.() -> Unit = {}, |
||||
|
) : AbstractTransport() { |
||||
|
private val TAG = "CustomSseClientTransport" |
||||
|
|
||||
|
private val scope by lazy { |
||||
|
CoroutineScope(session.coroutineContext + SupervisorJob()) |
||||
|
} |
||||
|
|
||||
|
// 使用Java标准库的AtomicBoolean替代kotlinx.atomicfu |
||||
|
private val initialized = AtomicBoolean(false) |
||||
|
private var session: ClientSSESession by Delegates.notNull() |
||||
|
private val endpoint = CompletableDeferred<String>() |
||||
|
|
||||
|
private var job: Job? = null |
||||
|
|
||||
|
// 创建JSON解析器,增强灵活性设置 |
||||
|
private val json = Json { |
||||
|
ignoreUnknownKeys = true // 忽略未知字段 |
||||
|
isLenient = true // 宽松解析模式 |
||||
|
coerceInputValues = true // 尝试强制转换类型 |
||||
|
encodeDefaults = true // 编码默认值 |
||||
|
explicitNulls = false // 不要求显式null值 |
||||
|
} |
||||
|
|
||||
|
// 保存基础URL(不包含查询参数)和查询参数 |
||||
|
private var baseUrlWithoutParams: String? = null |
||||
|
private var queryParams: Map<String, String> = emptyMap() |
||||
|
private var hostPart: String = "" // 添加类级别变量 |
||||
|
private var pathPart: String = "" // 添加类级别变量 |
||||
|
|
||||
|
/** |
||||
|
* 解析URL,分离基础URL、路径和查询参数 |
||||
|
* 返回三元组: (主机部分URL, 路径部分, 查询参数Map) |
||||
|
*/ |
||||
|
private fun parseUrl(url: String): Triple<String, String, Map<String, String>> { |
||||
|
return try { |
||||
|
val params = mutableMapOf<String, String>() |
||||
|
|
||||
|
// 确保URL有协议部分 |
||||
|
var processedUrl = url.trim() |
||||
|
if (!processedUrl.startsWith("http://") && !processedUrl.startsWith("https://")) { |
||||
|
processedUrl = "https://$processedUrl" |
||||
|
Log.d(TAG, "添加默认协议: $processedUrl") |
||||
|
} |
||||
|
|
||||
|
val urlObj = java.net.URL(processedUrl) |
||||
|
|
||||
|
// 解析查询参数 |
||||
|
if (urlObj.query != null) { |
||||
|
urlObj.query.split("&").forEach { param -> |
||||
|
val parts = param.split("=", limit = 2) |
||||
|
if (parts.size == 2) { |
||||
|
params[parts[0]] = parts[1] |
||||
|
} |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
// 构建主机部分URL(协议+主机+端口) |
||||
|
val port = if (urlObj.port == -1) "" else ":${urlObj.port}" |
||||
|
val hostUrl = "${urlObj.protocol}://${urlObj.host}$port" |
||||
|
|
||||
|
// 路径部分 |
||||
|
val path = urlObj.path |
||||
|
|
||||
|
Triple(hostUrl, path, params) |
||||
|
} catch (e: Exception) { |
||||
|
Log.e(TAG, "解析URL失败: $url, ${e.message}") |
||||
|
Triple(url, "", emptyMap()) |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
/** |
||||
|
* 收集SSE事件 |
||||
|
*/ |
||||
|
private suspend fun collectEvents() { |
||||
|
job = scope.launch(CoroutineName("CustomSseMcpClientTransport.collect#${hashCode()}")) { |
||||
|
session.incoming.collect { event -> |
||||
|
when (event.event) { |
||||
|
"error" -> { |
||||
|
val e = IllegalStateException("SSE error: ${event.data}") |
||||
|
Log.e(TAG, "SSE错误: ${event.data}") |
||||
|
_onError(e) |
||||
|
throw e |
||||
|
} |
||||
|
|
||||
|
"open" -> { |
||||
|
Log.d(TAG, "SSE连接已打开") |
||||
|
// 连接已打开,等待endpoint事件 |
||||
|
} |
||||
|
|
||||
|
"endpoint" -> { |
||||
|
try { |
||||
|
val eventData = event.data ?: "" |
||||
|
Log.d(TAG, "收到endpoint事件: $eventData") |
||||
|
|
||||
|
// 使用主机部分构建endpoint |
||||
|
val fullEndpoint = if (eventData.startsWith("/")) { |
||||
|
"$hostPart$eventData" |
||||
|
} else { |
||||
|
"$hostPart/$eventData" |
||||
|
} |
||||
|
|
||||
|
Log.d(TAG, "构建的endpoint路径(不含参数): $fullEndpoint") |
||||
|
|
||||
|
// 添加查询参数到endpoint |
||||
|
val endpointWithParams = if (queryParams.isNotEmpty()) { |
||||
|
// 检查endpoint是否已有查询参数 |
||||
|
if (fullEndpoint.contains("?")) { |
||||
|
// 已有查询参数,添加&并附加其他参数 |
||||
|
val queryString = queryParams.entries.joinToString("&") { "${it.key}=${it.value}" } |
||||
|
"$fullEndpoint&$queryString" |
||||
|
} else { |
||||
|
// 没有查询参数,添加?并附加参数 |
||||
|
val queryString = queryParams.entries.joinToString("&") { "${it.key}=${it.value}" } |
||||
|
"$fullEndpoint?$queryString" |
||||
|
} |
||||
|
} else { |
||||
|
fullEndpoint |
||||
|
} |
||||
|
|
||||
|
Log.d(TAG, "最终消息端点: $endpointWithParams") |
||||
|
endpoint.complete(endpointWithParams) |
||||
|
} catch (e: Exception) { |
||||
|
Log.e(TAG, "处理endpoint事件失败: ${e.message}", e) |
||||
|
_onError(e) |
||||
|
close() |
||||
|
error(e) |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
else -> { |
||||
|
try { |
||||
|
// 解析JSON-RPC消息 |
||||
|
val data = event.data |
||||
|
if (data != null) { |
||||
|
Log.d(TAG, "收到事件数据: $data") |
||||
|
try { |
||||
|
// 尝试安全地解析JSON消息 |
||||
|
safeParseMessage(data) |
||||
|
} catch (e: Exception) { |
||||
|
Log.e(TAG, "解析JSON-RPC消息失败: ${e.message}", e) |
||||
|
// 错误已记录,但不中断连接,只发送错误通知 |
||||
|
_onError(e) |
||||
|
} |
||||
|
} |
||||
|
} catch (e: Exception) { |
||||
|
Log.e(TAG, "处理事件失败: ${e.message}", e) |
||||
|
_onError(e) |
||||
|
} |
||||
|
} |
||||
|
} |
||||
|
} |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
/** |
||||
|
* 安全解析JSON-RPC消息 |
||||
|
*/ |
||||
|
private suspend fun safeParseMessage(data: String) { |
||||
|
try { |
||||
|
// 先尝试使用标准解析 |
||||
|
val message = json.decodeFromString<JSONRPCMessage>(data) |
||||
|
_onMessage(message) |
||||
|
} catch (e: Exception) { |
||||
|
// 如果标准解析失败,记录错误并尝试使用备用解析方式 |
||||
|
Log.w(TAG, "标准解析失败,尝试备用解析: ${e.message}") |
||||
|
|
||||
|
try { |
||||
|
// 尝试修复nextCursor缺失问题 |
||||
|
if (e.message?.contains("nextCursor") == true) { |
||||
|
// 尝试手动添加缺失的nextCursor字段 |
||||
|
val jsonObj = JSONObject(data) |
||||
|
|
||||
|
// 只有在解析ListToolsResult时处理 |
||||
|
if (data.contains("\"tools\"")) { |
||||
|
Log.d(TAG, "尝试修复ListToolsResult缺少nextCursor字段的问题") |
||||
|
|
||||
|
// 手动解析result部分并添加nextCursor |
||||
|
val resultJson = try { |
||||
|
if (jsonObj.has("result")) { |
||||
|
val resultObj = jsonObj.getJSONObject("result") |
||||
|
if (!resultObj.has("nextCursor")) { |
||||
|
resultObj.put("nextCursor", "") |
||||
|
jsonObj.put("result", resultObj) |
||||
|
} |
||||
|
jsonObj.toString() |
||||
|
} else { |
||||
|
// 如果没有result字段,可能是其他类型的消息 |
||||
|
data |
||||
|
} |
||||
|
} catch (ex: Exception) { |
||||
|
Log.e(TAG, "手动修复JSON失败: ${ex.message}") |
||||
|
data |
||||
|
} |
||||
|
|
||||
|
// 重新尝试解析修复后的JSON |
||||
|
val fixedMessage = json.decodeFromString<JSONRPCMessage>(resultJson) |
||||
|
_onMessage(fixedMessage) |
||||
|
return |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
// 通用错误处理 |
||||
|
Log.e(TAG, "无法解析消息,跳过: $data") |
||||
|
} catch (ex: Exception) { |
||||
|
Log.e(TAG, "备用解析也失败: ${ex.message}", ex) |
||||
|
// 不抛出异常,只记录错误 |
||||
|
_onError(e) |
||||
|
} |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
/** |
||||
|
* 启动传输层 |
||||
|
*/ |
||||
|
override suspend fun start() { |
||||
|
if (!initialized.compareAndSet(false, true)) { |
||||
|
Log.e(TAG, "传输层已经启动,不能重复启动") |
||||
|
error("CustomSseClientTransport already started!") |
||||
|
} |
||||
|
|
||||
|
// 解析URL和参数 |
||||
|
if (urlString != null) { |
||||
|
// 解析URL,提取主机部分、路径部分和查询参数 |
||||
|
val urlInfo = parseUrl(urlString) |
||||
|
hostPart = urlInfo.first |
||||
|
pathPart = urlInfo.second |
||||
|
queryParams = urlInfo.third |
||||
|
|
||||
|
// 存储不带查询参数的基础URL(主机+路径) |
||||
|
baseUrlWithoutParams = hostPart + pathPart |
||||
|
|
||||
|
Log.d(TAG, "原始URL: $urlString") |
||||
|
Log.d(TAG, "主机部分: $hostPart") |
||||
|
Log.d(TAG, "路径部分: $pathPart") |
||||
|
Log.d(TAG, "查询参数: $queryParams") |
||||
|
Log.d(TAG, "完整基础URL: $baseUrlWithoutParams") |
||||
|
} |
||||
|
|
||||
|
// 创建SSE会话 - 直接使用原始URL,不添加/sse后缀 |
||||
|
session = urlString?.let { |
||||
|
// 完整的SSE连接URL(主机部分+原始路径+查询参数) |
||||
|
val sseConnectUrl = if (queryParams.isNotEmpty()) { |
||||
|
// 如果路径已包含查询参数,就不再添加 |
||||
|
if (pathPart.contains("?")) { |
||||
|
"$hostPart$pathPart" |
||||
|
} else { |
||||
|
val queryString = queryParams.entries.joinToString("&") { "${it.key}=${it.value}" } |
||||
|
"$hostPart$pathPart?$queryString" |
||||
|
} |
||||
|
} else { |
||||
|
"$hostPart$pathPart" |
||||
|
} |
||||
|
|
||||
|
Log.d(TAG, "SSE连接URL: $sseConnectUrl") |
||||
|
|
||||
|
client.sseSession( |
||||
|
urlString = sseConnectUrl, |
||||
|
reconnectionTime = reconnectionTime, |
||||
|
block = requestBuilder, |
||||
|
) |
||||
|
} ?: client.sseSession( |
||||
|
reconnectionTime = reconnectionTime, |
||||
|
block = requestBuilder, |
||||
|
) |
||||
|
|
||||
|
// 收集SSE事件 |
||||
|
collectEvents() |
||||
|
|
||||
|
// 等待endpoint就绪 |
||||
|
endpoint.await() |
||||
|
Log.d(TAG, "传输层启动完成,消息端点已就绪") |
||||
|
} |
||||
|
|
||||
|
/** |
||||
|
* 发送消息 |
||||
|
*/ |
||||
|
@OptIn(ExperimentalCoroutinesApi::class) |
||||
|
override suspend fun send(message: JSONRPCMessage) { |
||||
|
if (!endpoint.isCompleted) { |
||||
|
Log.e(TAG, "发送失败: 未连接") |
||||
|
error("Not connected") |
||||
|
} |
||||
|
|
||||
|
try { |
||||
|
val messageEndpoint = endpoint.getCompleted() |
||||
|
Log.d(TAG, "发送消息到: $messageEndpoint") |
||||
|
|
||||
|
// 序列化消息 |
||||
|
val jsonString = json.encodeToString(message) |
||||
|
|
||||
|
val response = client.post(messageEndpoint) { |
||||
|
headers.append(HttpHeaders.ContentType, ContentType.Application.Json.toString()) |
||||
|
setBody(jsonString) |
||||
|
} |
||||
|
|
||||
|
if (!response.status.isSuccess()) { |
||||
|
val text = response.bodyAsText() |
||||
|
Log.e(TAG, "发送消息失败: HTTP ${response.status}, $text") |
||||
|
error("Error POSTing to endpoint (HTTP ${response.status}): $text") |
||||
|
} |
||||
|
} catch (e: Exception) { |
||||
|
Log.e(TAG, "发送消息异常: ${e.message}", e) |
||||
|
_onError(e) |
||||
|
throw e |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
/** |
||||
|
* 关闭传输层 |
||||
|
*/ |
||||
|
override suspend fun close() { |
||||
|
if (!initialized.get()) { |
||||
|
Log.e(TAG, "关闭失败: 传输层未初始化") |
||||
|
error("CustomSseClientTransport is not initialized!") |
||||
|
} |
||||
|
|
||||
|
session.cancel() |
||||
|
_onClose() |
||||
|
job?.cancelAndJoin() |
||||
|
Log.d(TAG, "传输层已关闭") |
||||
|
} |
||||
|
|
||||
|
/** |
||||
|
* 检查传输层是否已初始化 |
||||
|
*/ |
||||
|
fun isInitialized(): Boolean { |
||||
|
return initialized.get() |
||||
|
} |
||||
|
} |
||||
Loading…
Reference in new issue