Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@ import io.github.oshai.kotlinlogging.KotlinLogging
import io.modelcontextprotocol.kotlin.sdk.types.CancelledNotification
import io.modelcontextprotocol.kotlin.sdk.types.CancelledNotificationParams
import io.modelcontextprotocol.kotlin.sdk.types.EmptyResult
import io.modelcontextprotocol.kotlin.sdk.types.GetTaskPayloadResult
import io.modelcontextprotocol.kotlin.sdk.types.GetTaskResult
import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCEmptyMessage
import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCError
import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCNotification
Expand Down Expand Up @@ -51,7 +53,9 @@ import kotlinx.coroutines.withTimeoutOrNull
import kotlinx.serialization.SerializationException
import kotlinx.serialization.json.JsonObject
import kotlinx.serialization.json.JsonPrimitive
import kotlinx.serialization.json.decodeFromJsonElement
import kotlinx.serialization.json.encodeToJsonElement
import kotlinx.serialization.json.jsonObject
import kotlin.coroutines.AbstractCoroutineContextElement
import kotlin.coroutines.ContinuationInterceptor
import kotlin.coroutines.CoroutineContext
Expand Down Expand Up @@ -813,7 +817,7 @@ public abstract class Protocol(@PublishedApi internal val options: ProtocolOptio

try {
@Suppress("UNCHECKED_CAST")
result.complete(response!!.result as T)
result.complete(response!!.resultForRequest(request) as T)
} catch (e: Throwable) {
result.completeExceptionally(e)
}
Expand Down Expand Up @@ -901,6 +905,20 @@ public abstract class Protocol(@PublishedApi internal val options: ProtocolOptio
}
}

private fun JSONRPCResponse.resultForRequest(request: Request): RequestResult {
val rawResult = rawResultOrNull() ?: McpJson.encodeToJsonElement(result)
return when (request.method.value) {
Method.Defined.TasksGet.value,
Method.Defined.TasksCancel.value,
-> McpJson.decodeFromJsonElement<GetTaskResult>(rawResult)

Method.Defined.TasksResult.value,
-> GetTaskPayloadResult(rawResult.jsonObject)

else -> result
}
}

/**
* Emits a notification, which is a one-way message that does not expect a response.
*/
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ package io.modelcontextprotocol.kotlin.sdk.types
import kotlinx.serialization.EncodeDefault
import kotlinx.serialization.ExperimentalSerializationApi
import kotlinx.serialization.Serializable
import kotlinx.serialization.Transient
import kotlinx.serialization.json.JsonElement
import kotlinx.serialization.json.decodeFromJsonElement
import kotlinx.serialization.json.encodeToJsonElement
Expand Down Expand Up @@ -214,6 +215,15 @@ public data class JSONRPCResponse(val id: RequestId, val result: RequestResult =
/** Always `"2.0"` to indicate JSON-RPC 2.0 protocol. */
@EncodeDefault
override val jsonrpc: String = JSONRPC_VERSION

@Transient
private var rawResult: JsonElement? = null

internal fun withRawResult(result: JsonElement): JSONRPCResponse = apply {
rawResult = result
}

internal fun rawResultOrNull(): JsonElement? = rawResult
}

// ============================================================================
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ import kotlinx.serialization.json.JsonEncoder
import kotlinx.serialization.json.JsonObject
import kotlinx.serialization.json.JsonPrimitive
import kotlinx.serialization.json.contentOrNull
import kotlinx.serialization.json.decodeFromJsonElement
import kotlinx.serialization.json.jsonObject
import kotlinx.serialization.json.jsonPrimitive
import kotlinx.serialization.json.longOrNull
Expand Down Expand Up @@ -454,12 +455,15 @@ private fun selectServerResultDeserializer(element: JsonElement): Deserializatio
internal object RequestResultPolymorphicSerializer :
JsonContentPolymorphicSerializer<RequestResult>(RequestResult::class) {
override fun selectDeserializer(element: JsonElement): DeserializationStrategy<RequestResult> =
selectClientResultDeserializer(element)
?: selectServerResultDeserializer(element)
?: selectEmptyResult(element)
selectRequestResultDeserializer(element)
?: throw SerializationException("Cannot determine RequestResult type from JSON: ${element.jsonObject.keys}")
}

private fun selectRequestResultDeserializer(element: JsonElement): DeserializationStrategy<RequestResult>? =
selectClientResultDeserializer(element)
?: selectServerResultDeserializer(element)
?: selectEmptyResult(element)

/**
* Polymorphic serializer for [ClientResult] types.
* Throws [SerializationException] if the result type cannot be determined.
Expand Down Expand Up @@ -497,20 +501,64 @@ internal object ServerResultPolymorphicSerializer :
* - "method" + "id" -> JSONRPCRequest
* - "method" -> JSONRPCNotification
*/
internal object JSONRPCMessagePolymorphicSerializer :
JsonContentPolymorphicSerializer<JSONRPCMessage>(JSONRPCMessage::class) {
override fun selectDeserializer(element: JsonElement): DeserializationStrategy<JSONRPCMessage> {
internal object JSONRPCMessagePolymorphicSerializer : KSerializer<JSONRPCMessage> {
override val descriptor: SerialDescriptor = JSONRPCResponse.serializer().descriptor

override fun serialize(encoder: Encoder, value: JSONRPCMessage) {
when (value) {
is JSONRPCResponse -> encoder.encodeSerializableValue(JSONRPCResponse.serializer(), value)

is JSONRPCError -> encoder.encodeSerializableValue(JSONRPCError.serializer(), value)

is JSONRPCRequest -> encoder.encodeSerializableValue(JSONRPCRequest.serializer(), value)

is JSONRPCNotification -> encoder.encodeSerializableValue(JSONRPCNotification.serializer(), value)

JSONRPCEmptyMessage -> encoder.encodeSerializableValue(
JSONRPCEmptyMessage.serializer(),
JSONRPCEmptyMessage,
)
}
}

override fun deserialize(decoder: Decoder): JSONRPCMessage {
val jsonDecoder = decoder as? JsonDecoder
?: throw SerializationException("JSONRPCMessagePolymorphicSerializer requires a Json decoder")
val element = jsonDecoder.decodeJsonElement()
val jsonObj = element.jsonObject
return when {
"error" in jsonObj -> JSONRPCError.serializer()
"result" in jsonObj && "id" in jsonObj -> JSONRPCResponse.serializer()
"result" in jsonObj && jsonObj["result"]?.jsonObject?.isEmpty() == true -> JSONRPCEmptyMessage.serializer()
"method" in jsonObj && "id" in jsonObj -> JSONRPCRequest.serializer()
"method" in jsonObj -> JSONRPCNotification.serializer()
jsonObj.isEmpty() || jsonObj.keys == setOf("jsonrpc") -> JSONRPCEmptyMessage.serializer()
"error" in jsonObj -> jsonDecoder.json.decodeFromJsonElement(JSONRPCError.serializer(), element)

"result" in jsonObj && "id" in jsonObj -> decodeResponse(jsonDecoder, element)

"result" in jsonObj && jsonObj["result"]?.jsonObject?.isEmpty() == true ->
jsonDecoder.json.decodeFromJsonElement(JSONRPCEmptyMessage.serializer(), element)

"method" in jsonObj && "id" in jsonObj ->
jsonDecoder.json.decodeFromJsonElement(JSONRPCRequest.serializer(), element)

"method" in jsonObj -> jsonDecoder.json.decodeFromJsonElement(JSONRPCNotification.serializer(), element)

jsonObj.isEmpty() || jsonObj.keys == setOf("jsonrpc") ->
jsonDecoder.json.decodeFromJsonElement(JSONRPCEmptyMessage.serializer(), element)

else -> throw SerializationException("Invalid JSONRPCMessage type: ${jsonObj.keys}")
}
}

private fun decodeResponse(jsonDecoder: JsonDecoder, element: JsonElement): JSONRPCResponse {
val jsonObject = element.jsonObject
val rawResult = jsonObject.getValue("result")
val response = if (rawResult is JsonObject && selectRequestResultDeserializer(rawResult) == null) {
JSONRPCResponse(
id = jsonDecoder.json.decodeFromJsonElement(jsonObject.getValue("id")),
result = GetTaskPayloadResult(rawResult),
)
} else {
jsonDecoder.json.decodeFromJsonElement(JSONRPCResponse.serializer(), element)
}
return response.withRawResult(rawResult)
}
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,8 +8,17 @@ import io.kotest.matchers.shouldBe
import io.kotest.matchers.shouldNotBe
import io.kotest.matchers.string.shouldContain
import io.kotest.matchers.types.shouldBeInstanceOf
import io.modelcontextprotocol.kotlin.sdk.types.CallToolResult
import io.modelcontextprotocol.kotlin.sdk.types.CancelTaskRequest
import io.modelcontextprotocol.kotlin.sdk.types.CancelTaskRequestParams
import io.modelcontextprotocol.kotlin.sdk.types.CustomRequest
import io.modelcontextprotocol.kotlin.sdk.types.EmptyResult
import io.modelcontextprotocol.kotlin.sdk.types.GetTaskPayloadRequest
import io.modelcontextprotocol.kotlin.sdk.types.GetTaskPayloadRequestParams
import io.modelcontextprotocol.kotlin.sdk.types.GetTaskPayloadResult
import io.modelcontextprotocol.kotlin.sdk.types.GetTaskRequest
import io.modelcontextprotocol.kotlin.sdk.types.GetTaskRequestParams
import io.modelcontextprotocol.kotlin.sdk.types.GetTaskResult
import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCMessage
import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCNotification
import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCRequest
Expand All @@ -24,13 +33,17 @@ import io.modelcontextprotocol.kotlin.sdk.types.ReadResourceRequest
import io.modelcontextprotocol.kotlin.sdk.types.ReadResourceRequestParams
import io.modelcontextprotocol.kotlin.sdk.types.RequestId
import io.modelcontextprotocol.kotlin.sdk.types.RequestMeta
import io.modelcontextprotocol.kotlin.sdk.types.TaskStatus
import kotlinx.coroutines.async
import kotlinx.coroutines.launch
import kotlinx.coroutines.test.runTest
import kotlinx.serialization.json.JsonArray
import kotlinx.serialization.json.JsonElement
import kotlinx.serialization.json.JsonObject
import kotlinx.serialization.json.JsonObjectBuilder
import kotlinx.serialization.json.JsonPrimitive
import kotlinx.serialization.json.buildJsonObject
import kotlinx.serialization.json.decodeFromJsonElement
import kotlinx.serialization.json.encodeToJsonElement
import kotlinx.serialization.json.int
import kotlinx.serialization.json.jsonObject
Expand Down Expand Up @@ -195,6 +208,110 @@ class ProtocolTest {
inFlight.await()
}

@Test
fun `tasks result preserves payloads that overlap known result shapes`() = runTest {
protocol.connect(transport)
val payload = buildJsonObject {
put("content", JsonArray(emptyList()))
put("isError", JsonPrimitive(false))
put("extension", JsonPrimitive("preserved"))
}

val inFlight = async {
protocol.request<GetTaskPayloadResult>(
CustomRequest(method = Method.Custom(Method.Defined.TasksResult.value), params = null),
)
}

val sent = transport.awaitRequest()
transport.deliver(decodeResponse(sent.id, payload))

inFlight.await().json shouldBe payload
}

@Test
fun `tasks result preserves payloads that resemble task state`() = runTest {
protocol.connect(transport)
val payload = McpJson.encodeToJsonElement(completedTask("nested-task")).jsonObject

val inFlight = async {
protocol.request<GetTaskPayloadResult>(
GetTaskPayloadRequest(GetTaskPayloadRequestParams("task-1")),
)
}

val sent = transport.awaitRequest()
transport.deliver(decodeResponse(sent.id, payload))

inFlight.await().json shouldBe payload
}

@Test
fun `tasks result preserves otherwise unknown payloads`() = runTest {
protocol.connect(transport)
val payload = buildJsonObject {
put("extension", JsonPrimitive("preserved"))
}

val inFlight = async {
protocol.request<GetTaskPayloadResult>(
GetTaskPayloadRequest(GetTaskPayloadRequestParams("task-1")),
)
}

val sent = transport.awaitRequest()
transport.deliver(decodeResponse(sent.id, payload))

inFlight.await().json shouldBe payload
}

@Test
fun `tasks result adapts directly constructed transport responses`() = runTest {
protocol.connect(transport)
val responseResult = CallToolResult(content = emptyList(), isError = false)

val inFlight = async {
protocol.request<GetTaskPayloadResult>(
GetTaskPayloadRequest(GetTaskPayloadRequestParams("task-1")),
)
}

val sent = transport.awaitRequest()
transport.deliver(JSONRPCResponse(sent.id, responseResult))

inFlight.await().json shouldBe McpJson.encodeToJsonElement(responseResult).jsonObject
}

@Test
fun `tasks get preserves typed task state responses`() = runTest {
protocol.connect(transport)
val expected = completedTask("task-1")

val inFlight = async {
protocol.request<GetTaskResult>(GetTaskRequest(GetTaskRequestParams("task-1")))
}

val sent = transport.awaitRequest()
transport.deliver(decodeResponse(sent.id, McpJson.encodeToJsonElement(expected)))

inFlight.await() shouldBe expected
}

@Test
fun `tasks cancel preserves typed task state responses`() = runTest {
protocol.connect(transport)
val expected = completedTask("task-1").copy(status = TaskStatus.Cancelled)

val inFlight = async {
protocol.request<GetTaskResult>(CancelTaskRequest(CancelTaskRequestParams("task-1")))
}

val sent = transport.awaitRequest()
transport.deliver(decodeResponse(sent.id, McpJson.encodeToJsonElement(expected)))

inFlight.await() shouldBe expected
}

@Test
fun `request handler receives enriched extra and ambient context element`() = runTest {
protocol.connect(transport)
Expand Down Expand Up @@ -368,3 +485,19 @@ class ProtocolTest {
private fun metaOf(builderAction: JsonObjectBuilder.() -> Unit): RequestMeta = RequestMeta(metaJson(builderAction))

private fun metaJson(builderAction: JsonObjectBuilder.() -> Unit): JsonObject = buildJsonObject(builderAction)

private fun decodeResponse(id: RequestId, result: JsonElement): JSONRPCMessage = McpJson.decodeFromJsonElement(
buildJsonObject {
put("jsonrpc", JsonPrimitive("2.0"))
put("id", McpJson.encodeToJsonElement(id))
put("result", result)
},
)

private fun completedTask(taskId: String): GetTaskResult = GetTaskResult(
taskId = taskId,
status = TaskStatus.Completed,
createdAt = "2026-07-23T00:00:00Z",
lastUpdatedAt = "2026-07-23T00:00:01Z",
ttl = null,
)
Original file line number Diff line number Diff line change
Expand Up @@ -435,6 +435,25 @@ class JsonRpcTest {
)
}

@Test
fun `should round trip every JSONRPCMessage subtype polymorphically`() {
val messages = listOf<JSONRPCMessage>(
JSONRPCRequest(id = "request-1", method = "tools/list"),
JSONRPCNotification(method = "notifications/log"),
JSONRPCResponse(id = RequestId("response-1"), result = EmptyResult()),
JSONRPCError(
id = RequestId("error-1"),
error = RPCError(RPCError.ErrorCode.INTERNAL_ERROR, "Internal error"),
),
JSONRPCEmptyMessage,
)

messages.forEach { message ->
val encoded = McpJson.encodeToString<JSONRPCMessage>(message)
McpJson.decodeFromString<JSONRPCMessage>(encoded) shouldBe message
}
}

@Test
fun `should create JSONRPCRequest with string ID`() {
val params = buildJsonObject {
Expand Down
Loading