Sync OkHttpWebSocket with upstream quartz master

Port the relay-socket fixes from quartz's BasicOkHttpWebSocket into
Amber's fork while keeping the Tor/timeout-aware needsReconnect():

- Answer a relay-initiated CLOSE frame (onClosing -> close(1000))
  so OkHttp completes the handshake and reports onClosed at once
  instead of leaving the socket half-closed until the ping timeout.
- Guarantee a single terminal report per session via an
  AtomicBoolean claimed by onClosed/onFailure/disconnect, so late
  callbacks after disconnect no longer reach the relay client.
- disconnect() claims the session and synthesizes
  onClosed(1000, "client disconnect") synchronously, covering the
  case where OkHttp's cancel() raises no callback at all.
- connect() is idempotent (socket != null || ended -> return) and
  onOpen/onMessage drop stale events after the session ended.
- Mark socket/usingOkHttp @Volatile; they are written and read
  across OkHttp callback threads.

Add the upstream close-handshake regression test (raw RFC 6455
loopback relay, no new dependencies); both tests fail on the
previous implementation.
This commit is contained in:
greenart7c3
2026-09-14 09:33:13 -03:00
parent d1ca5a277f
commit 7d44c39b34
2 changed files with 320 additions and 23 deletions
@@ -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
@@ -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<Throwable?>(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()
}
}
}