diff --git a/app/src/main/java/com/greenart7c3/nostrsigner/okhttp/OkHttpWebSocket.kt b/app/src/main/java/com/greenart7c3/nostrsigner/okhttp/OkHttpWebSocket.kt index c1e5fcac..e64758aa 100644 --- a/app/src/main/java/com/greenart7c3/nostrsigner/okhttp/OkHttpWebSocket.kt +++ b/app/src/main/java/com/greenart7c3/nostrsigner/okhttp/OkHttpWebSocket.kt @@ -25,6 +25,7 @@ import com.vitorpamplona.quartz.nip01Core.relay.sockets.WebSocket import com.vitorpamplona.quartz.nip01Core.relay.sockets.WebSocketListener import com.vitorpamplona.quartz.nip01Core.relay.sockets.WebsocketBuilder import com.vitorpamplona.quartz.nip01Core.relay.sockets.okhttp.BasicOkHttpWebSocket.Companion.exceptionHandler +import java.util.concurrent.atomic.AtomicBoolean import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.cancel @@ -40,8 +41,23 @@ class OkHttpWebSocket( val httpClient: (url: NormalizedRelayUrl) -> OkHttpClient, val out: WebSocketListener, ) : WebSocket { - private var usingOkHttp: OkHttpClient? = null - private var socket: okhttp3.WebSocket? = null + @Volatile private var usingOkHttp: OkHttpClient? = null + + @Volatile private var socket: okhttp3.WebSocket? = null + + /** + * Set once, by whichever of `onClosed`, `onFailure` or [disconnect] ends the session first. + * + * One adapter is one session: the relay client builds a fresh one per dial, and OkHttp binds + * exactly one socket to the listener created in [connect], so anything that reaches that + * listener is from this session by construction. The only question a callback has to ask is + * whether the session already ended -- which is what keeps the [WebSocket.disconnect] contract: + * after [disconnect] the failure OkHttp raises for its own `cancel()` on the reader thread, or + * the `onClosed` its writer thread delivers once a close handshake completes, is dropped rather + * than reaching a relay client that has already moved on. Claimed with a compare-and-set so a + * [disconnect] racing a terminal callback still yields exactly one report. + */ + private val ended = AtomicBoolean(false) fun buildRequest() = Request.Builder().url(url.url).build() @@ -67,6 +83,8 @@ class OkHttpWebSocket( } override fun connect() { + if (socket != null || ended.get()) return + usingOkHttp = httpClient(url) socket = usingOkHttp?.newWebSocket(buildRequest(), OkHttpWebsocketListener(out)) } @@ -90,35 +108,65 @@ class OkHttpWebSocket( } } + /** Claims the session's single terminal report. False if it already ended. */ + private fun endSession(): Boolean { + if (!ended.compareAndSet(false, true)) return false + socket = null + incomingMessages.close() + job.cancel() + scope.cancel() + return true + } + override fun onOpen( webSocket: okhttp3.WebSocket, response: Response, - ) = out.onOpen( - (response.receivedResponseAtMillis - response.sentRequestAtMillis).toInt(), - response.headers["Sec-WebSocket-Extensions"]?.contains("permessage-deflate") ?: false, - ) + ) { + if (ended.get()) return + out.onOpen( + (response.receivedResponseAtMillis - response.sentRequestAtMillis).toInt(), + response.headers["Sec-WebSocket-Extensions"]?.contains("permessage-deflate") ?: false, + ) + } override fun onMessage( webSocket: okhttp3.WebSocket, text: String, ) { - // Asynchronously send the received message to the channel. - // `trySendBlocking` is used here for simplicity within the callback, - // but it's important to understand potential thread blocking if the buffer is full. + if (ended.get()) return + // Never blocks (unlimited channel): the OkHttp reader + // thread must stay free to keep draining the socket. incomingMessages.trySendBlocking(text) } + override fun onClosing( + webSocket: okhttp3.WebSocket, + code: Int, + reason: String, + ) { + // The relay sent a CLOSE frame. OkHttp's contract (WebSocketListener KDoc, + // RealWebSocket, and its own WebSocketEcho recipe) is that onClosed fires + // only once BOTH peers have sent a close, and sending ours is the + // application's job. Left unanswered, the socket sits half-closed: no + // onClosed, no onFailure, send() still accepted and silently discarded, and + // a later cancel() is silent too -- so the relay client kept believing it + // was connected, with its REQs live, until OkHttp's 120s ping path finally + // failed up to two intervals later. Answering completes the handshake and + // OkHttp reports onClosed at once, whether or not the relay still holds the + // TCP session open. + // + // Always 1000 rather than echoing `code`: close() validates the code it is + // asked to write and throws on the reserved ones (1005, 1006, 1015), and a + // relay may send anything. + webSocket.close(1000, null) + } + override fun onClosed( webSocket: okhttp3.WebSocket, code: Int, reason: String, ) { - // Close the channel when the WebSocket connection is closed. - incomingMessages.close() - job.cancel() - scope.cancel() - - socket = null + if (!endSession()) return out.onClosed(code, reason) } @@ -127,12 +175,7 @@ class OkHttpWebSocket( t: Throwable, response: Response?, ) { - // Close the channel on failure, and propagate the error. - incomingMessages.close() - job.cancel() - scope.cancel() - - socket = null + if (!endSession()) return out.onFailure(t, response?.code, response?.message) } } @@ -148,9 +191,14 @@ class OkHttpWebSocket( } override fun disconnect() { - // uses cancel to kill the SEND stack that might be waiting - socket?.cancel() + // Claim the session ourselves: OkHttp's cancel() raises no callback when no reader is + // left to fail (the state a relay-initiated close leaves behind), and when it does the + // failure arrives later on its own thread. The relay client needs the answer now. + val closing = socket ?: return + if (!ended.compareAndSet(false, true)) return socket = null + closing.cancel() + out.onClosed(1000, "client disconnect") } override fun send(msg: String): Boolean = socket?.send(msg) ?: false diff --git a/app/src/test/java/com/greenart7c3/nostrsigner/okhttp/OkHttpWebSocketCloseHandshakeTest.kt b/app/src/test/java/com/greenart7c3/nostrsigner/okhttp/OkHttpWebSocketCloseHandshakeTest.kt new file mode 100644 index 00000000..cbbb4c3c --- /dev/null +++ b/app/src/test/java/com/greenart7c3/nostrsigner/okhttp/OkHttpWebSocketCloseHandshakeTest.kt @@ -0,0 +1,249 @@ +/* + * Copyright (c) 2025 Vitor Pamplona + * + * Permission is hereby granted, free of charge, to any person obtaining a copy of + * this software and associated documentation files (the "Software"), to deal in + * the Software without restriction, including without limitation the rights to use, + * copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the + * Software, and to permit persons to whom the Software is furnished to do so, + * subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS + * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR + * COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN + * AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION + * WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + */ +package com.greenart7c3.nostrsigner.okhttp + +import com.vitorpamplona.quartz.nip01Core.relay.normalizer.NormalizedRelayUrl +import com.vitorpamplona.quartz.nip01Core.relay.sockets.WebSocketListener +import java.io.InputStream +import java.io.OutputStream +import java.net.ServerSocket +import java.net.Socket +import java.security.MessageDigest +import java.util.Base64 +import java.util.concurrent.CountDownLatch +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.atomic.AtomicReference +import kotlin.concurrent.thread +import okhttp3.OkHttpClient +import org.junit.Assert.assertEquals +import org.junit.Assert.assertNull +import org.junit.Assert.assertTrue +import org.junit.Test + +/** + * A relay-initiated close must be answered, or OkHttp never finishes the handshake. + * + * OkHttp fires `onClosed` only once BOTH peers have sent a CLOSE frame, and sending ours is the + * application's job (its `WebSocketEcho` recipe answers `onClosing` with `close(1000, null)`). + * Before [OkHttpWebSocket] did that, a relay's CLOSE frame left the socket half-closed: no + * `onClosed`, no `onFailure`, `send()` still accepted and discarded, and a later `cancel()` silent + * too. The relay client kept believing it was connected, with its REQs live, until the ping path + * failed up to two ping intervals later. + * + * Driven against a minimal RFC 6455 server on a loopback [ServerSocket] rather than a mock + * server library, so the test needs no new dependency and controls the exact frames on the wire. + */ +class OkHttpWebSocketCloseHandshakeTest { + /** Handshakes one client, sends it a CLOSE frame on demand, and records the frames it sends back. */ + private class TinyRelay : AutoCloseable { + private val server = ServerSocket(0) + val url = NormalizedRelayUrl(url = "ws://127.0.0.1:${server.localPort}/") + + private val handshaken = CountDownLatch(1) + val clientCloseFrame = CountDownLatch(1) + val clientCloseCode = AtomicInteger(-1) + + private var socket: Socket? = null + private var out: OutputStream? = null + + private val thread = + thread(isDaemon = true, name = "tiny-relay") { + runCatching { + val s = server.accept() + socket = s + val input = s.getInputStream() + val output = s.getOutputStream() + out = output + handshake(input, output) + handshaken.countDown() + readFrames(input) + } + } + + private fun handshake( + input: InputStream, + output: OutputStream, + ) { + var key: String? = null + val line = StringBuilder() + while (true) { + val c = input.read() + check(c != -1) { "EOF during handshake" } + if (c == '\n'.code) { + val l = line.toString().trim() + if (l.isEmpty()) break + if (l.lowercase().startsWith("sec-websocket-key:")) key = l.substring(18).trim() + line.setLength(0) + } else if (c != '\r'.code) { + line.append(c.toChar()) + } + } + val accept = + Base64.getEncoder().encodeToString( + MessageDigest.getInstance("SHA-1").digest((key + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11").toByteArray()), + ) + output.write( + ( + "HTTP/1.1 101 Switching Protocols\r\n" + + "Upgrade: websocket\r\n" + + "Connection: Upgrade\r\n" + + "Sec-WebSocket-Accept: $accept\r\n\r\n" + ).toByteArray(Charsets.ISO_8859_1), + ) + output.flush() + } + + /** Client frames are masked; decode enough to spot a CLOSE and read its status code. */ + private fun readFrames(input: InputStream) { + while (true) { + val b0 = input.read() + if (b0 == -1) return + val b1 = input.read() + if (b1 == -1) return + val opcode = b0 and 0x0F + var len = b1 and 0x7F + if (len == 126) { + len = (input.read() shl 8) or input.read() + } else if (len == 127) { + len = 0 + repeat(8) { len = (len shl 8) or input.read() } + } + val masked = (b1 and 0x80) != 0 + val mask = if (masked) ByteArray(4) { input.read().toByte() } else ByteArray(4) + val payload = ByteArray(len) { i -> (input.read() xor mask[i % 4].toInt()).toByte() } + if (opcode == 0x8) { + if (len >= 2) { + clientCloseCode.set(((payload[0].toInt() and 0xFF) shl 8) or (payload[1].toInt() and 0xFF)) + } + clientCloseFrame.countDown() + } + } + } + + fun awaitClient() = handshaken.await(5, TimeUnit.SECONDS) + + /** Server-initiated CLOSE, status 1000, unmasked as servers send it. The TCP session stays open. */ + fun sendClose() { + val output = checkNotNull(out) { "no client yet" } + output.write(byteArrayOf(0x88.toByte(), 0x02, 0x03, 0xE8.toByte())) + output.flush() + } + + override fun close() { + runCatching { socket?.close() } + runCatching { server.close() } + thread.join(2_000) + } + } + + private class Recorder : WebSocketListener { + val opened = CountDownLatch(1) + val closed = CountDownLatch(1) + val closedCount = AtomicInteger(0) + val closedCode = AtomicInteger(-1) + val failure = AtomicReference(null) + + override fun onOpen( + pingMillis: Int, + compression: Boolean, + ) = opened.countDown() + + override suspend fun onMessage(text: String) {} + + override fun onClosed( + code: Int, + reason: String, + ) { + closedCode.set(code) + closedCount.incrementAndGet() + closed.countDown() + } + + override fun onFailure( + t: Throwable, + code: Int?, + response: String?, + ) { + failure.set(t) + } + } + + @Test + fun `a relay initiated close is answered and reported as closed`() { + TinyRelay().use { relay -> + val recorder = Recorder() + val client = OkHttpClient() + val socket = OkHttpWebSocket(relay.url, { client }, recorder) + + socket.connect() + assertTrue("relay never saw the client", relay.awaitClient()) + assertTrue("no onOpen", recorder.opened.await(5, TimeUnit.SECONDS)) + + relay.sendClose() + + // The half of the handshake that is ours to send. + assertTrue("client never answered the relay's CLOSE frame", relay.clientCloseFrame.await(5, TimeUnit.SECONDS)) + assertEquals(1000, relay.clientCloseCode.get()) + + // And the terminal callback the relay client's bookkeeping depends on. + assertTrue("onClosed never fired", recorder.closed.await(5, TimeUnit.SECONDS)) + assertEquals("the relay's status code is what gets reported", 1000, recorder.closedCode.get()) + assertNull("a clean handshake is not a failure", recorder.failure.get()) + + // The session already ended; the usual teardown afterwards must not report it twice. + socket.disconnect() + assertEquals("one terminal report per session", 1, recorder.closedCount.get()) + assertTrue("a closed socket needs a fresh dial", socket.needsReconnect()) + + client.dispatcher.executorService.shutdown() + } + } + + @Test + fun `disconnect reports the session end once, synchronously, and drops what OkHttp says afterwards`() { + TinyRelay().use { relay -> + val recorder = Recorder() + val client = OkHttpClient() + val socket = OkHttpWebSocket(relay.url, { client }, recorder) + + socket.connect() + assertTrue("relay never saw the client", relay.awaitClient()) + assertTrue("no onOpen", recorder.opened.await(5, TimeUnit.SECONDS)) + + socket.disconnect() + + // Reported before disconnect() returned: the relay client dials the replacement + // right after this call and must not hear from the old socket later. + assertEquals("disconnect() must report synchronously", 1, recorder.closedCount.get()) + assertEquals(1000, recorder.closedCode.get()) + assertTrue(socket.needsReconnect()) + + // OkHttp's own reaction to cancel() -- a failure on its reader thread -- and the relay's + // reaction to the dropped TCP session must both be swallowed. + Thread.sleep(500) + assertEquals("no second report", 1, recorder.closedCount.get()) + assertNull("the cancel's failure must not surface", recorder.failure.get()) + + client.dispatcher.executorService.shutdown() + } + } +}