Skip to content
Merged
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
@@ -0,0 +1,66 @@
package io.tolgee.websocket

import org.springframework.web.socket.CloseStatus
import org.springframework.web.socket.WebSocketMessage
import org.springframework.web.socket.WebSocketSession
import org.springframework.web.socket.handler.ConcurrentWebSocketSessionDecorator
import java.util.concurrent.TimeUnit
import java.util.concurrent.locks.ReentrantLock
import kotlin.concurrent.withLock

/**
* Spring writes a STOMP ERROR frame and closes with PROTOCOL_ERROR straight after it. If another thread holds
* the flush lock at that moment the frame is only buffered, and the close makes that thread discard it — the
* client sees a bare disconnect instead of the `Unauthenticated` ERROR the webapp stops reconnecting on.
* Remove once https://github.com/spring-projects/spring-framework/issues/37328 is fixed.
*/
class ErrorFrameFlushingSessionDecorator(
delegate: WebSocketSession,
sendTimeLimit: Int,
bufferSizeLimit: Int,
) : ConcurrentWebSocketSessionDecorator(delegate, sendTimeLimit, bufferSizeLimit) {
private var closing = false
private var sendsInProgress = 0
private val sendsLock = ReentrantLock()
private val allSendsFinished = sendsLock.newCondition()

override fun sendMessage(message: WebSocketMessage<*>) {
sendsLock.withLock {
if (closing) return
sendsInProgress++
}
try {
super.sendMessage(message)
} finally {
sendsLock.withLock {
if (--sendsInProgress == 0) allSendsFinished.signalAll()
}
}
}

override fun close(status: CloseStatus) {
sendsLock.withLock { closing = true }
if (status.equalsCode(CloseStatus.PROTOCOL_ERROR)) {
awaitSendsInProgress()
}
super.close(status)
}

/**
* A thread leaving `sendMessage` has flushed the buffer unless its write failed, and after a failed write
* nothing is left to send the rest, so only sends in progress are worth waiting for.
*/
private fun awaitSendsInProgress() {
var remainingNanos = TimeUnit.MILLISECONDS.toNanos((sendTimeLimit - timeSinceSendStarted).coerceAtLeast(0))
sendsLock.withLock {
while (sendsInProgress > 0 && remainingNanos > 0) {
try {
remainingNanos = allSendsFinished.awaitNanos(remainingNanos)
} catch (e: InterruptedException) {
Thread.currentThread().interrupt()
return
}
}
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
package io.tolgee.websocket

import org.springframework.context.annotation.Bean
import org.springframework.context.annotation.Configuration
import org.springframework.messaging.support.AbstractSubscribableChannel
import org.springframework.web.socket.WebSocketHandler
import org.springframework.web.socket.WebSocketSession
import org.springframework.web.socket.config.annotation.DelegatingWebSocketMessageBrokerConfiguration
import org.springframework.web.socket.messaging.SubProtocolWebSocketHandler

/**
* Takes the place of `@EnableWebSocketMessageBroker` (which only imports the parent class) so the session
* decorator can be swapped for [ErrorFrameFlushingSessionDecorator]. Putting the annotation back drops that fix.
*/
@Configuration(proxyBeanMethods = false)
class WebSocketBrokerConfiguration : DelegatingWebSocketMessageBrokerConfiguration() {
@Bean
override fun subProtocolWebSocketHandler(
clientInboundChannel: AbstractSubscribableChannel,
clientOutboundChannel: AbstractSubscribableChannel,
): WebSocketHandler =
object : SubProtocolWebSocketHandler(clientInboundChannel, clientOutboundChannel) {
override fun decorateSession(session: WebSocketSession): WebSocketSession =
ErrorFrameFlushingSessionDecorator(session, sendTimeLimit, sendBufferSizeLimit)
}.also { it.phase = phase }
}
Original file line number Diff line number Diff line change
Expand Up @@ -15,14 +15,12 @@ import org.springframework.messaging.simp.stomp.StompCommand
import org.springframework.messaging.simp.stomp.StompHeaderAccessor
import org.springframework.messaging.support.ChannelInterceptor
import org.springframework.messaging.support.MessageHeaderAccessor
import org.springframework.web.socket.config.annotation.EnableWebSocketMessageBroker
import org.springframework.web.socket.config.annotation.StompEndpointRegistry
import org.springframework.web.socket.config.annotation.WebSocketMessageBrokerConfigurer
import java.security.Principal

/** Websocket authentication and per-topic authorization model: docs/websocket/README.md. */
@Configuration
@EnableWebSocketMessageBroker
class WebSocketConfig(
@Lazy
private val websocketAuthenticationResolver: WebsocketAuthenticationResolver,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
package io.tolgee.websocket

import io.tolgee.testing.assert
import org.junit.jupiter.api.Test
import org.mockito.kotlin.any
import org.mockito.kotlin.doAnswer
import org.mockito.kotlin.doReturn
import org.mockito.kotlin.mock
import org.springframework.web.socket.CloseStatus
import org.springframework.web.socket.TextMessage
import org.springframework.web.socket.WebSocketMessage
import org.springframework.web.socket.WebSocketSession
import java.io.IOException
import java.util.concurrent.CopyOnWriteArrayList
import java.util.concurrent.CountDownLatch
import kotlin.concurrent.thread
import kotlin.system.measureTimeMillis

class ErrorFrameFlushingSessionDecoratorTest {
private val written = CopyOnWriteArrayList<String>()
private val sendEntered = CountDownLatch(1)
private val releaseSend = CountDownLatch(1)

private fun delegate(failBlockedSend: Boolean = false) =
mock<WebSocketSession> {
on { id } doReturn "session"
on { isOpen } doReturn true
on { sendMessage(any()) } doAnswer {
val payload = (it.arguments[0] as WebSocketMessage<*>).payload.toString()
if (payload == BLOCKED_FRAME) {
sendEntered.countDown()
releaseSend.await()
if (failBlockedSend) throw IOException("client gone")
}
written.add(payload)
Unit
}
on { close(any()) } doAnswer {
written.add(CLOSED)
Unit
}
}

@Test
fun `a protocol-error close writes the buffered ERROR frame first and drops frames sent after it`() {
val session = ErrorFrameFlushingSessionDecorator(delegate(), SEND_TIME_LIMIT_MS, BUFFER_LIMIT)
val sender = thread { session.sendMessage(TextMessage(BLOCKED_FRAME)) }
sendEntered.await()
session.sendMessage(TextMessage("ERROR"))

val closer = thread { session.close(CloseStatus.PROTOCOL_ERROR) }
awaitWaiting(closer)
session.sendMessage(TextMessage("MESSAGE after the ERROR"))
releaseSend.countDown()
closer.join()
sender.join()

written.assert.containsExactly(BLOCKED_FRAME, "ERROR", CLOSED)
}

@Test
fun `a protocol-error close does not wait for frames nobody is left to send`() {
val session =
ErrorFrameFlushingSessionDecorator(delegate(failBlockedSend = true), SEND_TIME_LIMIT_MS, BUFFER_LIMIT)
val sender =
thread {
runCatching { session.sendMessage(TextMessage(BLOCKED_FRAME)) }
}
sendEntered.await()
session.sendMessage(TextMessage("ERROR"))
releaseSend.countDown()
sender.join()

val closeTookMs = measureTimeMillis { session.close(CloseStatus.PROTOCOL_ERROR) }

closeTookMs.assert.isLessThan(SEND_TIME_LIMIT_MS / 2L)
}

@Test
fun `a protocol-error close stops waiting once the active send has overrun the send time limit`() {
val session = ErrorFrameFlushingSessionDecorator(delegate(), SEND_TIME_LIMIT_MS, BUFFER_LIMIT)
val stuckSender = thread { session.sendMessage(TextMessage(BLOCKED_FRAME)) }
try {
sendEntered.await()
session.sendMessage(TextMessage("ERROR"))
Thread.sleep(SEND_TIME_LIMIT_MS + 100L)

val closeTookMs = measureTimeMillis { session.close(CloseStatus.PROTOCOL_ERROR) }

closeTookMs.assert.isLessThan(SEND_TIME_LIMIT_MS / 2L)
} finally {
releaseSend.countDown()
stuckSender.join()
}
}

private fun awaitWaiting(thread: Thread) {
while (thread.state != Thread.State.TIMED_WAITING && thread.state != Thread.State.WAITING) {
Thread.onSpinWait()
}
}

companion object {
private const val SEND_TIME_LIMIT_MS = 1000
private const val BUFFER_LIMIT = 1024 * 1024
private const val BLOCKED_FRAME = "CONNECTED"
private const val CLOSED = "<closed>"
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
package io.tolgee.websocket

import io.tolgee.AbstractSpringTest
import io.tolgee.testing.WebsocketTest
import io.tolgee.websocket.WebsocketTestHelper.Auth
import org.junit.jupiter.api.Test
import org.springframework.boot.test.context.SpringBootTest
import org.springframework.boot.test.context.TestConfiguration
import org.springframework.boot.test.web.server.LocalServerPort
import org.springframework.context.annotation.Import
import org.springframework.web.socket.CloseStatus
import org.springframework.web.socket.WebSocketMessage
import org.springframework.web.socket.WebSocketSession
import org.springframework.web.socket.config.annotation.WebSocketMessageBrokerConfigurer
import org.springframework.web.socket.config.annotation.WebSocketTransportRegistration
import org.springframework.web.socket.handler.WebSocketHandlerDecorator
import org.springframework.web.socket.handler.WebSocketSessionDecorator
import java.util.concurrent.ConcurrentHashMap
import java.util.concurrent.CountDownLatch
import java.util.concurrent.TimeUnit

@SpringBootTest(
properties = ["tolgee.websocket.use-redis=false"],
webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT,
)
@Import(WebsocketErrorFrameDeliveryTest.HoldFlushLockConfiguration::class)
@WebsocketTest
class WebsocketErrorFrameDeliveryTest : AbstractSpringTest() {
@LocalServerPort
private val port: Int? = null

@Test
fun `the Unauthenticated ERROR frame reaches the client while CONNECTED is still being flushed`() {
val socket = WebsocketTestHelper(port, Auth(jwtToken = "invalid"), projectId = 1, userId = 1)
try {
socket.listenForTranslationDataModified()
socket.waitForUnauthenticated()
} finally {
socket.stop()
}
}

/**
* Keeps the session's flush lock held after CONNECTED is on the wire — the way an outbound thread
* descheduled on a loaded CI runner does — until the rejected SUBSCRIBE has arrived and the close that
* follows the ERROR frame is requested. Releasing it on the SUBSCRIBE alone lets the ERROR go out
* uncontended, and the test then passes even when the frame can be lost.
*/
@TestConfiguration
class HoldFlushLockConfiguration : WebSocketMessageBrokerConfigurer {
override fun configureWebSocketTransport(registration: WebSocketTransportRegistration) {
registration.addDecoratorFactory { handler ->
object : WebSocketHandlerDecorator(handler) {
override fun afterConnectionEstablished(session: WebSocketSession) {
super.afterConnectionEstablished(HoldAfterConnectedSession(session))
}

override fun handleMessage(
session: WebSocketSession,
message: WebSocketMessage<*>,
) {
if (message.isStompFrame("SUBSCRIBE")) {
latchesFor(session).subscribeReceived.countDown()
}
super.handleMessage(session, message)
}
}
}
}
}

private class HoldAfterConnectedSession(
delegate: WebSocketSession,
) : WebSocketSessionDecorator(delegate) {
override fun sendMessage(message: WebSocketMessage<*>) {
super.sendMessage(message)
if (message.isStompFrame("CONNECTED")) {
val latches = latchesFor(this)
latches.subscribeReceived.await(SUBSCRIBE_TIMEOUT_MS, TimeUnit.MILLISECONDS)
latches.closeRequested.await(CLOSE_TIMEOUT_MS, TimeUnit.MILLISECONDS)
}
}

override fun close(status: CloseStatus) {
latchesFor(this).closeRequested.countDown()
super.close(status)
}
}

private class SessionLatches {
val subscribeReceived = CountDownLatch(1)
val closeRequested = CountDownLatch(1)
}

companion object {
private const val SUBSCRIBE_TIMEOUT_MS = 5000L

private const val CLOSE_TIMEOUT_MS = 1000L

private val latchesBySessionId = ConcurrentHashMap<String, SessionLatches>()

private fun latchesFor(session: WebSocketSession) =
latchesBySessionId.computeIfAbsent(session.id) { SessionLatches() }

private fun WebSocketMessage<*>.isStompFrame(command: String) = (payload as? String)?.startsWith(command) == true
}
}
Loading