Skip to content
Open
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
Expand Up @@ -17,6 +17,7 @@

package org.connectbot.sshlib.client

import kotlinx.coroutines.CancellationException
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.channels.Channel
import kotlinx.coroutines.channels.ReceiveChannel
Expand Down Expand Up @@ -61,6 +62,13 @@ internal class ForwardingChannel(
connection.sendWindowAdjust(remoteChannelNumber, adjust)
}
}
} catch (cancelled: CancellationException) {
throw cancelled
} catch (failure: Exception) {
// Writes admitted before cancellation may fail with a transport error instead of
// CancellationException. Surface it to readers and the connection, never globally.
_incomingData.close(failure)
connection.transportFailed(failure)
} finally {
_incomingData.close()
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@

package org.connectbot.sshlib.client

import kotlinx.coroutines.CancellationException
import kotlinx.coroutines.CompletableDeferred
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.CoroutineStart
Expand Down Expand Up @@ -88,6 +89,10 @@ class SessionChannel internal constructor(
_stdout.send(buffer.toByteArray())
stdoutBufferConsumed.trySend(Unit)
}
} catch (cancelled: CancellationException) {
throw cancelled
} catch (failure: Exception) {
_stdout.close(failure)
} finally {
_stdout.close()
}
Expand Down Expand Up @@ -215,6 +220,11 @@ class SessionChannel internal constructor(
connection.sendWindowAdjust(_remoteChannelNumber, adjust)
}
}
} catch (cancelled: CancellationException) {
throw cancelled
} catch (failure: Exception) {
output.close(failure)
connection.transportFailed(failure)
} finally {
output.close()
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -358,7 +358,7 @@ class SshConnection(
internal val connectionScope = CoroutineScope(SupervisorJob() + coroutineDispatcher)
private val protocolScope = CoroutineScope(SupervisorJob() + stateMachineDispatcher)
internal val protocolExecutor = ProtocolExecutor(protocolScope, stateMachineDispatcher)
private val outboundPacketController = PacketWriter(connectionScope, stateMachineDispatcher, packetIO, ::writerFailed)
private val outboundPacketController = PacketWriter(connectionScope, stateMachineDispatcher, packetIO, ::transportFailed)
private val closeMutex = Mutex()
private var transportClosing = false

Expand Down Expand Up @@ -521,7 +521,8 @@ class SshConnection(
outboundPacketController.writePacket(messageType, payload)
}

private suspend fun writerFailed(failure: Throwable) {
/** Background channel delivery has no caller to receive failed window-update writes. */
internal suspend fun transportFailed(failure: Throwable) {
if (transportClosing || !connectionScope.isActive) return
protocolScope.launch {
_disconnectedFlow.tryEmit(failure)
Expand All @@ -538,6 +539,10 @@ class SshConnection(
transportClosing = true
try {
transport.close()
} catch (failure: Exception) {
// A nested SSH transport may send CHANNEL_CLOSE on an upstream that has
// already disconnected. Cleanup must still terminate our own protocol loop.
logger.debug("Transport close failed", failure)
} finally {
outboundPacketController.close()
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.test.UnconfinedTestDispatcher
import kotlinx.coroutines.test.runTest
import org.connectbot.sshlib.SshException
import org.connectbot.sshlib.transport.TransportException
import org.junit.jupiter.api.Assertions.assertArrayEquals
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertFalse
Expand All @@ -33,6 +34,20 @@ import kotlin.test.assertFailsWith

class ForwardingChannelTest {

@Test
fun `failed window update closes delivery with its cause instead of escaping the worker`() = runTest {
val conn = mockk<SshConnection>(relaxed = true)
val failure = TransportException("Transport closed")
coEvery { conn.sendWindowAdjust(any(), any()) } throws failure
val (channel, _) = createChannel(connection = conn, initialWindowSize = 128)

channel.onData(ByteArray(100))
assertEquals(100, channel.incomingData.receive().size)

assertEquals(failure.message, channel.incomingData.receiveCatching().exceptionOrNull()?.message)
coVerify(exactly = 1) { conn.transportFailed(failure) }
}

private fun createChannel(
connection: SshConnection = mockk(relaxed = true),
remoteWindowSize: Long = 64 * 1024,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ import kotlinx.coroutines.test.UnconfinedTestDispatcher
import kotlinx.coroutines.test.runCurrent
import kotlinx.coroutines.test.runTest
import org.connectbot.sshlib.SshException
import org.connectbot.sshlib.transport.TransportException
import org.junit.jupiter.api.Assertions.assertArrayEquals
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertFalse
Expand All @@ -40,6 +41,23 @@ import kotlin.test.assertFailsWith
@OptIn(ExperimentalCoroutinesApi::class)
class SessionChannelTest {

@Test
fun `failed window update closes stdout with its cause instead of escaping the worker`() = runTest {
for (buffered in listOf(false, true)) {
val conn = mockk<SshConnection>(relaxed = true)
val failure = TransportException("Transport closed")
coEvery { conn.sendWindowAdjust(any(), any()) } throws failure
val (channel, _) = createChannel(connection = conn, initialWindowSize = 128, bufferedStdout = buffered)

channel.onData(ByteArray(100))
assertEquals(100, channel.stdout.receive().size)

assertEquals(failure.message, channel.stdout.receiveCatching().exceptionOrNull()?.message)
coVerify(exactly = 1) { conn.transportFailed(failure) }
channel.close()
}
}

private fun createChannel(
connection: SshConnection = mockk(relaxed = true),
initialWindowSize: Int = 64 * 1024,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,19 @@ import kotlin.test.assertFailsWith
@OptIn(ExperimentalCoroutinesApi::class)
class SshConnectionCloseTest {

@Test
fun `failing nested transport close still completes shutdown exactly once`() = runTest {
val transport = RecordingTransport(failClose = true)
val connection = connection(transport, StandardTestDispatcher(testScheduler))

connection.close()
connection.close()

assertEquals(1, transport.closeCalls)
assertFailsWith<TransportException> { connection.sendChannelClose(recipientChannel = 0) }
assertEquals(0, transport.writeCalls)
}

@Test
fun `concurrent close closes transport once`() = runTest {
val transport = RecordingTransport(suspendClose = true)
Expand Down Expand Up @@ -74,6 +87,7 @@ class SshConnectionCloseTest {

private class RecordingTransport(
private val suspendClose: Boolean = false,
private val failClose: Boolean = false,
) : Transport {
val closeStarted = CompletableDeferred<Unit>()
val allowClose = CompletableDeferred<Unit>()
Expand All @@ -93,6 +107,7 @@ class SshConnectionCloseTest {
closeCalls++
closeStarted.complete(Unit)
if (suspendClose) allowClose.await()
if (failClose) throw TransportException("Upstream packet writer stopped")
}
}
}
Loading