From 1fecc81a71b656cf774e7baa40c9e4ee020c2877 Mon Sep 17 00:00:00 2001 From: taekop Date: Wed, 23 Sep 2026 10:51:18 +0900 Subject: [PATCH] fix(server): run onInitialized callbacks registered after initialization (#920) ServerSession.onInitialized composed callbacks into a plain var, so a callback registered after notifications/initialized was never invoked, and concurrent registration could drop callbacks. Keep pending callbacks in an AtomicRef. The initialized handler swaps it to null and runs the list in registration order; registration after that runs the callback immediately. Each callback now runs at most once. Fixes #920 --- .../sdk/server/ServerSessionInitializeTest.kt | 47 +++++++++++++++++++ .../kotlin/sdk/server/ServerSession.kt | 22 ++++++--- 2 files changed, 62 insertions(+), 7 deletions(-) diff --git a/integration-test/src/jvmTest/kotlin/io/modelcontextprotocol/kotlin/sdk/server/ServerSessionInitializeTest.kt b/integration-test/src/jvmTest/kotlin/io/modelcontextprotocol/kotlin/sdk/server/ServerSessionInitializeTest.kt index 38b169ae9..a97b71803 100644 --- a/integration-test/src/jvmTest/kotlin/io/modelcontextprotocol/kotlin/sdk/server/ServerSessionInitializeTest.kt +++ b/integration-test/src/jvmTest/kotlin/io/modelcontextprotocol/kotlin/sdk/server/ServerSessionInitializeTest.kt @@ -5,6 +5,7 @@ import io.modelcontextprotocol.kotlin.sdk.types.ClientCapabilities import io.modelcontextprotocol.kotlin.sdk.types.Implementation import io.modelcontextprotocol.kotlin.sdk.types.InitializeRequest import io.modelcontextprotocol.kotlin.sdk.types.InitializeRequestParams +import io.modelcontextprotocol.kotlin.sdk.types.InitializedNotification import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCError import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCMessage import io.modelcontextprotocol.kotlin.sdk.types.JSONRPCRequest @@ -25,6 +26,7 @@ import kotlinx.serialization.json.put import kotlinx.serialization.json.putJsonObject import org.junit.jupiter.api.Test import java.util.concurrent.CopyOnWriteArrayList +import java.util.concurrent.atomic.AtomicInteger import kotlin.test.assertEquals import kotlin.test.assertFalse import kotlin.test.assertNotNull @@ -191,4 +193,49 @@ class ServerSessionInitializeTest { assertEquals(RPCError.ErrorCode.INVALID_REQUEST, error.error.code) } } + + private suspend fun completeHandshake(session: ServerSession) { + val (clientTransport, serverTransport) = InMemoryTransport.createLinkedPair() + val responseDone = CompletableDeferred() + clientTransport.onMessage { message -> + if (message is JSONRPCResponse) responseDone.complete(message) + } + + session.connect(serverTransport) + clientTransport.send(createInitializeRequest().toJSON()) + responseDone.await() + clientTransport.send(InitializedNotification().toJSON()) + } + + @Test + fun `should run onInitialized callback registered after initialization`() = runTest { + val session = createSession() + val calls = CopyOnWriteArrayList() + val initialized = CompletableDeferred() + + session.onInitialized { calls.add("early") } + session.onInitialized { initialized.complete(Unit) } + + completeHandshake(session) + initialized.await() + + session.onInitialized { calls.add("late") } + + assertEquals(listOf("early", "late"), calls) + } + + @Test + fun `should run every onInitialized callback registered concurrently`() = runTest { + val session = createSession() + val counters = List(200) { AtomicInteger() } + + withContext(Dispatchers.Default) { + counters.map { counter -> + launch { session.onInitialized { counter.incrementAndGet() } } + }.joinAll() + } + completeHandshake(session) + + assertEquals(List(counters.size) { 1 }, counters.map { it.get() }) + } } diff --git a/kotlin-sdk-server/src/commonMain/kotlin/io/modelcontextprotocol/kotlin/sdk/server/ServerSession.kt b/kotlin-sdk-server/src/commonMain/kotlin/io/modelcontextprotocol/kotlin/sdk/server/ServerSession.kt index f20d3e122..94b3f71a7 100644 --- a/kotlin-sdk-server/src/commonMain/kotlin/io/modelcontextprotocol/kotlin/sdk/server/ServerSession.kt +++ b/kotlin-sdk-server/src/commonMain/kotlin/io/modelcontextprotocol/kotlin/sdk/server/ServerSession.kt @@ -31,6 +31,9 @@ import io.modelcontextprotocol.kotlin.sdk.types.SUPPORTED_PROTOCOL_VERSIONS import io.modelcontextprotocol.kotlin.sdk.types.SetLevelRequest import kotlinx.atomicfu.AtomicRef import kotlinx.atomicfu.atomic +import kotlinx.atomicfu.getAndUpdate +import kotlinx.collections.immutable.PersistentList +import kotlinx.collections.immutable.persistentListOf import kotlinx.coroutines.CompletableDeferred import kotlinx.serialization.json.JsonObject import kotlin.uuid.ExperimentalUuidApi @@ -64,7 +67,12 @@ public open class ServerSession( @OptIn(ExperimentalUuidApi::class) public val sessionId: String = Uuid.random().toString() - private var _onInitialized: (() -> Unit) = {} + /** + * Callbacks waiting for `notifications/initialized`, or `null` once it has been received. + * A callback registered after that point runs immediately instead of being stored. + */ + private val pendingInitializedCallbacks: AtomicRef Unit>?> = + atomic(persistentListOf()) private var _onClose: () -> Unit = {} @@ -94,7 +102,7 @@ public open class ServerSession( handleInitialize(request) } setNotificationHandler(Defined.NotificationsInitialized) { - _onInitialized() + pendingInitializedCallbacks.getAndSet(null)?.forEach { callback -> callback() } CompletableDeferred(Unit) } @@ -120,13 +128,13 @@ public open class ServerSession( * * The callback must be synchronous and fast: it runs on the message-dispatch path for * `notifications/initialized`, after concurrent dispatch has been enabled for the session. + * + * If initialization has already completed, [block] runs immediately on the calling thread. + * Otherwise it runs when `notifications/initialized` is received, after the callbacks + * registered before it. Each callback runs at most once. */ public fun onInitialized(block: () -> Unit) { - val old = _onInitialized - _onInitialized = { - old() - block() - } + if (pendingInitializedCallbacks.getAndUpdate { it?.adding(block) } == null) block() } /**