From 673265d785d2e71733194d13c873401915efe803 Mon Sep 17 00:00:00 2001 From: toki Date: Sun, 26 Apr 2026 07:06:15 +0900 Subject: [PATCH] feat: add TlsWsTest for Kotlin TLS WebSocket testing --- .../com/tokilabs/toki_socket/TlsWsTest.kt | 108 ++++++++++++++++++ 1 file changed, 108 insertions(+) create mode 100644 kotlin/src/test/kotlin/com/tokilabs/toki_socket/TlsWsTest.kt diff --git a/kotlin/src/test/kotlin/com/tokilabs/toki_socket/TlsWsTest.kt b/kotlin/src/test/kotlin/com/tokilabs/toki_socket/TlsWsTest.kt new file mode 100644 index 0000000..e4d4176 --- /dev/null +++ b/kotlin/src/test/kotlin/com/tokilabs/toki_socket/TlsWsTest.kt @@ -0,0 +1,108 @@ +package com.tokilabs.toki_socket + +import com.tokilabs.toki_socket.packets.TestData +import kotlinx.coroutines.CompletableDeferred +import kotlinx.coroutines.delay +import kotlinx.coroutines.runBlocking +import java.io.IOException +import java.nio.file.Paths +import kotlin.test.Test +import kotlin.test.assertEquals + +class TlsWsTest { + @Test + fun testWssSendReceive() = runBlocking { + val (serverContext, clientContext) = createTlsContexts() + val received = CompletableDeferred() + val listenerReady = CompletableDeferred() + val server = WsServer("127.0.0.1", 0, "/", sslContext = serverContext) { conn -> + WsClient.forServer(conn, 0, 0, testParserMap()) + } + server.onClientConnected = { client -> + addListenerTyped(client.communicator) { received.complete(it) } + listenerReady.complete(Unit) + } + var client: WsClient? = null + server.start() + try { + client = dialWssWithRetry("127.0.0.1", server.port(), "/", clientContext) + + listenerReady.await() + client.send(TestData.newBuilder().setIndex(55).setMessage("hello wss").build()) + + assertEquals(55, received.await().index) + } finally { + closeClientAndServer(client, server) + } + } + + @Test + fun testWssRequestResponse() = runBlocking { + val (serverContext, clientContext) = createTlsContexts() + val server = WsServer("127.0.0.1", 0, "/", sslContext = serverContext) { conn -> + WsClient.forServer(conn, 0, 0, testParserMap()) + } + server.onClientConnected = { client -> + addRequestListenerTyped(client.communicator) { req -> + TestData.newBuilder() + .setIndex(req.index * 2) + .setMessage("wss echo: ${req.message}") + .build() + } + } + var client: WsClient? = null + server.start() + try { + client = dialWssWithRetry("127.0.0.1", server.port(), "/", clientContext) + + val res = sendRequestTyped( + client.communicator, + TestData.newBuilder().setIndex(11).setMessage("wss req").build(), + timeoutMs = 2_000, + ) + + assertEquals(22, res.index) + assertEquals("wss echo: wss req", res.message) + } finally { + closeClientAndServer(client, server) + } + } + + private fun createTlsContexts(): Pair { + val loader = Thread.currentThread().contextClassLoader + val certPath = Paths.get(loader.getResource("server.crt")!!.toURI()).toString() + val keyPath = Paths.get(loader.getResource("server.key")!!.toURI()).toString() + return createTestSslContexts(certPath, keyPath) + } + + private suspend fun dialWssWithRetry( + host: String, + port: Int, + path: String, + sslContext: javax.net.ssl.SSLContext, + ): WsClient { + var lastError: IOException? = null + repeat(3) { attempt -> + try { + return dialWss(host, port, path, sslContext, 0, 0, testParserMap()) + } catch (ex: IOException) { + lastError = ex + if (attempt < 2) delay(100) + } + } + throw lastError ?: IOException("wss dial failed") + } + + private suspend fun closeClientAndServer(client: WsClient?, server: WsServer) { + client?.close() + try { + if (client != null) { + waitForCondition(message = "wss client did not disconnect") { + server.clients().isEmpty() + } + } + } finally { + server.stop() + } + } +}