diff --git a/src/Renci.SshNet/Session.cs b/src/Renci.SshNet/Session.cs index 5468b53c8..a5d662779 100644 --- a/src/Renci.SshNet/Session.cs +++ b/src/Renci.SshNet/Session.cs @@ -1045,17 +1045,41 @@ internal void WaitOnHandle(WaitHandle waitHandle, TimeSpan timeout) /// The size of the packet exceeds the maximum size defined by the protocol. internal void SendMessage(Message message) { - if (!_socket.IsConnected()) + while (true) { - throw new SshConnectionException("Client not connected."); - } + if (!_socket.IsConnected()) + { + throw new SshConnectionException("Client not connected."); + } - if (!_keyExchangeCompletedWaitHandle.IsSet && message is not IKeyExchangedAllowed) - { - // Wait for key exchange to be completed - WaitOnHandle(_keyExchangeCompletedWaitHandle.WaitHandle); + if (!_keyExchangeCompletedWaitHandle.IsSet && message is not IKeyExchangedAllowed) + { + // Wait for key exchange to be completed + WaitOnHandle(_keyExchangeCompletedWaitHandle.WaitHandle); + } + + // take a write lock to ensure the outbound packet sequence number is incremented + // atomically, and only after the packet has actually been sent + lock (_socketWriteLock) + { + if (!_keyExchangeCompletedWaitHandle.IsSet && message is not IKeyExchangedAllowed) + { + // A key re-exchange started between the check above and acquiring the + // write lock. Our SSH_MSG_KEXINIT may already have been sent, in which + // case sending this message now would violate RFC 4253 section 7.1 and + // cause the server to drop the connection. Go back to waiting for the + // key exchange to complete. + continue; + } + + SendMessageWithinWriteLock(message); + return; + } } + } + private void SendMessageWithinWriteLock(Message message) + { if (_logger.IsEnabled(LogLevel.Trace)) { _logger.LogTrace("[{SessionId}] Sending message {MessageName}({MessageNumber}) to server: '{Message}'.", SessionIdHex, message.MessageName, message.MessageNumber, message.ToString()); @@ -1074,82 +1098,77 @@ internal void SendMessage(Message message) macLength = _clientMac.HashSize / 8; } - // take a write lock to ensure the outbound packet sequence number is incremented - // atomically, and only after the packet has actually been sent - lock (_socketWriteLock) - { - var activeBufferLength = message.GetPacket( - ref _sendBuffer, - paddingMultiplier, - _clientCompression, - _clientEtm || _clientAead, - macLength); + var activeBufferLength = message.GetPacket( + ref _sendBuffer, + paddingMultiplier, + _clientCompression, + _clientEtm || _clientAead, + macLength); - // write outbound packet sequence to start of packet data - BinaryPrimitives.WriteUInt32BigEndian(_sendBuffer, _outboundPacketSequence); + // write outbound packet sequence to start of packet data + BinaryPrimitives.WriteUInt32BigEndian(_sendBuffer, _outboundPacketSequence); - if (_clientMac != null && !_clientEtm) - { - // non-ETM mac = MAC(key, sequence_number || unencrypted_packet) + if (_clientMac != null && !_clientEtm) + { + // non-ETM mac = MAC(key, sequence_number || unencrypted_packet) - var hashSuccess = _clientMac.TryComputeHash( - buffer: _sendBuffer, - offset: 0, - count: activeBufferLength - macLength, - destination: _sendBuffer.AsSpan(activeBufferLength - macLength), - bytesWritten: out var bytesWritten); + var hashSuccess = _clientMac.TryComputeHash( + buffer: _sendBuffer, + offset: 0, + count: activeBufferLength - macLength, + destination: _sendBuffer.AsSpan(activeBufferLength - macLength), + bytesWritten: out var bytesWritten); - Debug.Assert(hashSuccess && bytesWritten == macLength); - } + Debug.Assert(hashSuccess && bytesWritten == macLength); + } - if (_clientCipher != null) - { - _clientCipher.SetSequenceNumber(_outboundPacketSequence); + if (_clientCipher != null) + { + _clientCipher.SetSequenceNumber(_outboundPacketSequence); - // Not encrypting the sequence number (it is not part of the packet), - // nor the packet length for ETM. - var offset = _clientEtm ? 8 : 4; + // Not encrypting the sequence number (it is not part of the packet), + // nor the packet length for ETM. + var offset = _clientEtm ? 8 : 4; - var numberOfBytesEncrypted = _clientCipher.Encrypt( - input: _sendBuffer, - offset, - length: activeBufferLength - offset - macLength, - output: _sendBuffer, - outputOffset: offset); + var numberOfBytesEncrypted = _clientCipher.Encrypt( + input: _sendBuffer, + offset, + length: activeBufferLength - offset - macLength, + output: _sendBuffer, + outputOffset: offset); - Debug.Assert(numberOfBytesEncrypted == activeBufferLength - offset - macLength + (_clientAead ? macLength : 0)); - } + Debug.Assert(numberOfBytesEncrypted == activeBufferLength - offset - macLength + (_clientAead ? macLength : 0)); + } - if (_clientMac != null && _clientEtm) - { - // ETM mac = MAC(key, sequence_number || packet_length || encrypted_packet) + if (_clientMac != null && _clientEtm) + { + // ETM mac = MAC(key, sequence_number || packet_length || encrypted_packet) - var hashSuccess = _clientMac.TryComputeHash( - buffer: _sendBuffer, - offset: 0, - count: activeBufferLength - macLength, - destination: _sendBuffer.AsSpan(activeBufferLength - macLength), - bytesWritten: out var bytesWritten); + var hashSuccess = _clientMac.TryComputeHash( + buffer: _sendBuffer, + offset: 0, + count: activeBufferLength - macLength, + destination: _sendBuffer.AsSpan(activeBufferLength - macLength), + bytesWritten: out var bytesWritten); - Debug.Assert(hashSuccess && bytesWritten == macLength); - } + Debug.Assert(hashSuccess && bytesWritten == macLength); + } - SendPacket(_sendBuffer, 4, activeBufferLength - 4); + SendPacket(_sendBuffer, 4, activeBufferLength - 4); - if (_isStrictKex && message is NewKeysMessage) - { - _outboundPacketSequence = 0; - } - else - { - // increment the packet sequence number only after we're sure the packet has - // been sent; even though it's only used for the MAC, it needs to be incremented - // for each package sent. - // - // the server will use it to verify the data integrity, and as such the order in - // which messages are sent must follow the outbound packet sequence number - _outboundPacketSequence++; - } + if (_isStrictKex && message is NewKeysMessage) + { + _outboundPacketSequence = 0; + } + else + { + // increment the packet sequence number only after we're sure the packet has + // been sent; even though it's only used for the MAC, it needs to be incremented + // for each package sent. + // + // the server will use it to verify the data integrity, and as such the order in + // which messages are sent must follow the outbound packet sequence number + _outboundPacketSequence++; } } diff --git a/test/Renci.SshNet.IntegrationTests/.dockerignore b/test/Renci.SshNet.IntegrationTests/.dockerignore index a8eb4de7e..5f3d7f5e8 100644 --- a/test/Renci.SshNet.IntegrationTests/.dockerignore +++ b/test/Renci.SshNet.IntegrationTests/.dockerignore @@ -1,3 +1,4 @@ * !server/* !user/* +!proftpd/* diff --git a/test/Renci.SshNet.IntegrationTests/Logging/TextWriterLogger.cs b/test/Renci.SshNet.IntegrationTests/Logging/TextWriterLogger.cs new file mode 100644 index 000000000..43694f89d --- /dev/null +++ b/test/Renci.SshNet.IntegrationTests/Logging/TextWriterLogger.cs @@ -0,0 +1,45 @@ +#nullable enable + +using Microsoft.Extensions.Logging; + +namespace Renci.SshNet.IntegrationTests.Logging +{ + internal class TextWriterLogger(TextWriter writer, string categoryName) : ILogger + { + public IDisposable? BeginScope(TState state) + where TState : notnull + { + return null; + } + + public bool IsEnabled(LogLevel logLevel) + { + return true; + } + + public void Log(LogLevel logLevel, EventId eventId, TState state, Exception? exception, Func formatter) + { + StringBuilder sb = new StringBuilder(); + sb.Append(logLevel); + sb.Append(": "); + sb.Append(categoryName); + sb.Append(": "); + + string message = formatter(state, exception); + sb.Append(message); + + if (exception != null) + { + sb.Append(": "); + sb.Append(exception); + } + + string line = sb.ToString(); + + lock (writer) + { + writer.WriteLine(line); + } + } + } +} diff --git a/test/Renci.SshNet.IntegrationTests/Logging/TextWriterLoggerProvider.cs b/test/Renci.SshNet.IntegrationTests/Logging/TextWriterLoggerProvider.cs new file mode 100644 index 000000000..ea1715b0b --- /dev/null +++ b/test/Renci.SshNet.IntegrationTests/Logging/TextWriterLoggerProvider.cs @@ -0,0 +1,18 @@ +#nullable enable + +using Microsoft.Extensions.Logging; + +namespace Renci.SshNet.IntegrationTests.Logging +{ + internal class TextWriterLoggerProvider(TextWriter writer) : ILoggerProvider + { + public ILogger CreateLogger(string categoryName) + { + return new TextWriterLogger(writer, categoryName); + } + + public void Dispose() + { + } + } +} diff --git a/test/Renci.SshNet.IntegrationTests/ProFtpdRekeyTests.cs b/test/Renci.SshNet.IntegrationTests/ProFtpdRekeyTests.cs new file mode 100644 index 000000000..a66e9f70a --- /dev/null +++ b/test/Renci.SshNet.IntegrationTests/ProFtpdRekeyTests.cs @@ -0,0 +1,164 @@ +#if NET // The test uses Parallel.ForEachAsync, which is not available on .NET Framework. + +using DotNet.Testcontainers.Builders; +using DotNet.Testcontainers.Containers; +using DotNet.Testcontainers.Images; + +using Microsoft.Extensions.Logging; + +using Renci.SshNet.IntegrationTests.Logging; + +namespace Renci.SshNet.IntegrationTests +{ + /// + /// Reproduces https://github.com/sshnet/SSH.NET/issues/1764: connection failures during + /// SFTP transfers when the server initiates a key re-exchange. + /// + /// Unlike OpenSSH, which queues non key exchange output while a re-exchange is in progress, + /// ProFTPD mod_sftp keeps sending channel messages (SSH_MSG_CHANNEL_WINDOW_ADJUST, + /// SSH_MSG_CHANNEL_DATA) after it has sent its SSH_MSG_KEXINIT. The same load which passes + /// against the OpenSSH test server therefore fails against ProFTPD with either + /// "Message type 93 is not valid in the current context." or a connection drop + /// ("Key exchange failed"). These tests run against a ProFTPD container configured to + /// re-key every 1 MB (see proftpd/proftpd.conf) to give the race many trials per upload. + /// + /// + /// Trace-level logging to a file is enabled for the duration of these tests: the small + /// per-message overhead on the message listener thread widens the window between the + /// arrival of the server's SSH_MSG_KEXINIT and its processing, during which concurrent + /// uploaders keep sending data. This mirrors real-world conditions (applications with + /// trace logging enabled, or slower links) and makes the race fail reliably on loopback. + /// + /// + [TestClass] + public sealed class ProFtpdRekeyTests : TestBase + { + private static IFutureDockerImage _proFtpdImage; + private static IContainer _proFtpdServer; + private static string _proFtpdHostName; + private static ushort _proFtpdPort; + private static StreamWriter _traceLogWriter; + private static ILoggerFactory _traceLoggerFactory; + + [ClassInitialize] + [System.Diagnostics.CodeAnalysis.SuppressMessage("Style", "IDE0060:Remove unused parameter", Justification = "MSTests requires context parameter")] + public static async Task ClassInitialize(TestContext context) + { + // The Windows Tests in CI cannot run the ProFTPD container: Docker on the Windows + // runners is in Windows containers mode ("no matching manifest for windows/amd64"), + // which is why the OpenSSH server for the other integration tests is set up in + // WSL2 with Podman instead (see InfrastructureFixture). + if (OperatingSystem.IsWindows() && Environment.GetEnvironmentVariable("CI") == "true") + { + Assert.Inconclusive("Requires a container runtime able to run Linux containers."); + } + + _traceLogWriter = new StreamWriter(Path.GetTempFileName()) { AutoFlush = true }; + _traceLoggerFactory = LoggerFactory.Create(builder => + { + builder.SetMinimumLevel(LogLevel.Trace); + builder.AddProvider(new TextWriterLoggerProvider(_traceLogWriter)); + }); + + SshNetLoggingConfiguration.InitializeLogging(_traceLoggerFactory); + + _proFtpdImage = new ImageFromDockerfileBuilder() + .WithName("renci-ssh-tests-proftpd-image") + .WithDockerfileDirectory(CommonDirectoryPath.GetSolutionDirectory(), Path.Combine("test", "Renci.SshNet.IntegrationTests")) + .WithDockerfile("proftpd/Dockerfile") + .WithDeleteIfExists(true) + .Build(); + + await _proFtpdImage.CreateAsync(context.CancellationToken); + + _proFtpdServer = new ContainerBuilder(_proFtpdImage) + .WithHostname("renci-ssh-tests-proftpd") + .WithPortBinding(22, true) + .WithWaitStrategy(Wait.ForUnixContainer().UntilInternalTcpPortIsAvailable(22)) + .Build(); + + await _proFtpdServer.StartAsync(context.CancellationToken); + + _proFtpdPort = _proFtpdServer.GetMappedPublicPort(22); + _proFtpdHostName = _proFtpdServer.Hostname; + } + + [ClassCleanup] + public static async Task ClassCleanup() + { + if (_proFtpdServer != null) + { + await _proFtpdServer.DisposeAsync(); + } + + if (_proFtpdImage != null) + { + await _proFtpdImage.DisposeAsync(); + } + + // Restore the assembly-wide logging configuration set up by InfrastructureFixture. + var defaultLoggerFactory = LoggerFactory.Create(builder => + { + builder.SetMinimumLevel(LogLevel.Debug); + builder.AddTestConsoleLogger(); + }); + + SshNetLoggingConfiguration.InitializeLogging(defaultLoggerFactory); + + _traceLoggerFactory?.Dispose(); + _traceLogWriter?.Dispose(); + } + + [TestMethod] + public async Task Sftp_ConcurrentUploads_WithServerRekey() + { + const int fileSize = 128 * 1024 * 1024; + const int concurrentUploads = 4; + const int attempts = 3; + + using (var sftp = new SftpClient(_proFtpdHostName, _proFtpdPort, "sshnet", "ssh4ever")) + { + await sftp.ConnectAsync(CancellationToken.None); + + for (var attempt = 0; attempt < attempts; attempt++) + { + await Parallel.ForEachAsync(Enumerable.Range(0, concurrentUploads), async (i, ct) => + { + var localFile = CreateZeroFilledTempFile(fileSize); + + try + { + var remoteFile = $"rekey-test-{i}"; + + using (var fileStream = File.OpenRead(localFile)) + { + await sftp.UploadFileAsync(fileStream, remoteFile, ct); + } + + var remoteLength = (await sftp.GetAsync(remoteFile, ct)).Attributes.Size; + + Assert.AreEqual(fileSize, remoteLength); + } + finally + { + File.Delete(localFile); + } + }); + } + } + } + + private static string CreateZeroFilledTempFile(long size) + { + var file = Path.GetTempFileName(); + + using (var fs = File.OpenWrite(file)) + { + fs.SetLength(size); + } + + return file; + } + } +} +#endif diff --git a/test/Renci.SshNet.IntegrationTests/proftpd/Dockerfile b/test/Renci.SshNet.IntegrationTests/proftpd/Dockerfile new file mode 100644 index 000000000..138e7c358 --- /dev/null +++ b/test/Renci.SshNet.IntegrationTests/proftpd/Dockerfile @@ -0,0 +1,11 @@ +# ProFTPD with mod_sftp, used by ProFtpdRekeyTests to reproduce +# https://github.com/sshnet/SSH.NET/issues/1764 (alpine:3.24 ships ProFTPD 1.3.9c). +FROM alpine:3.24 +RUN apk add --no-cache proftpd proftpd-mod_sftp openssh-keygen +RUN ssh-keygen -t rsa -b 3072 -N "" -m PEM -f /etc/proftpd/host_rsa +RUN chmod 400 /etc/proftpd/host_rsa +RUN adduser -D sshnet +RUN echo 'sshnet:ssh4ever' | chpasswd +COPY proftpd/proftpd.conf /etc/proftpd/proftpd.conf +EXPOSE 22 +CMD ["proftpd", "--nodaemon", "--config", "/etc/proftpd/proftpd.conf"] diff --git a/test/Renci.SshNet.IntegrationTests/proftpd/proftpd.conf b/test/Renci.SshNet.IntegrationTests/proftpd/proftpd.conf new file mode 100644 index 000000000..8bd8ee6d7 --- /dev/null +++ b/test/Renci.SshNet.IntegrationTests/proftpd/proftpd.conf @@ -0,0 +1,26 @@ +ServerName "renci-ssh-tests-proftpd" +ServerType standalone +DefaultServer on +# SFTPEngine converts this server to SSH/SFTP entirely; no plain FTP is served. +Port 22 +User nobody +Group nobody +DefaultRoot ~ +AllowOverwrite on +PidFile /var/run/proftpd.pid +ScoreboardFile /var/run/proftpd.scoreboard + +LoadModule mod_sftp.c + + + SFTPEngine on + SFTPHostKey /etc/proftpd/host_rsa + SFTPAuthMethods password + SFTPLog /var/log/sftp.log + # Re-key every 1 MB so that the key re-exchange races of + # https://github.com/sshnet/SSH.NET/issues/1764 get many trials per upload. + # Unlike OpenSSH, ProFTPD mod_sftp keeps sending channel messages + # (SSH_MSG_CHANNEL_WINDOW_ADJUST / SSH_MSG_CHANNEL_DATA) after it has + # initiated a re-exchange, which is what triggers the client-side failures. + SFTPRekey required 3600 1 +