diff --git a/src/libraries/System.Net.Security/src/System/Net/Security/SslStream.IO.cs b/src/libraries/System.Net.Security/src/System/Net/Security/SslStream.IO.cs index bc79976d479d5b..53794a221d17d6 100644 --- a/src/libraries/System.Net.Security/src/System/Net/Security/SslStream.IO.cs +++ b/src/libraries/System.Net.Security/src/System/Net/Security/SslStream.IO.cs @@ -455,6 +455,10 @@ private async ValueTask ReceiveHandshakeFrameAsync(Cancellation throw new IOException(SR.net_io_eof); } +#pragma warning disable CS0618 + int handshakeTypeOffset = _lastFrame.Header.Version == SslProtocols.Ssl2 ? HandshakeTypeOffsetSsl2 : HandshakeTypeOffsetTls; +#pragma warning restore CS0618 + // At this point, we have at least one TLS frame. switch (_lastFrame.Header.Type) { @@ -465,10 +469,13 @@ private async ValueTask ReceiveHandshakeFrameAsync(Cancellation } break; case TlsContentType.Handshake: -#pragma warning disable CS0618 - if (!_isRenego && _buffer.EncryptedReadOnlySpan[_lastFrame.Header.Version == SslProtocols.Ssl2 ? HandshakeTypeOffsetSsl2 : HandshakeTypeOffsetTls] == (byte)TlsHandshakeType.ClientHello && + if (frameSize <= handshakeTypeOffset) + { + throw new IOException(SR.net_ssl_io_frame); + } + + if (!_isRenego && _buffer.EncryptedReadOnlySpan[handshakeTypeOffset] == (byte)TlsHandshakeType.ClientHello && _sslAuthenticationOptions!.IsServer) // guard against malicious endpoints. We should not see ClientHello on client. -#pragma warning restore CS0618 { TlsFrameHelper.ProcessingOptions options = TlsFrameHelper.ProcessingOptions.ServerName; @@ -535,10 +542,8 @@ private async ValueTask ReceiveHandshakeFrameAsync(Cancellation // (_securityContext == null), reject any frame that is not a ClientHello. if (_sslAuthenticationOptions!.IsServer && _securityContext == null) { -#pragma warning disable CS0618 bool isClientHello = _lastFrame.Header.Type == TlsContentType.Handshake && - _buffer.EncryptedReadOnlySpan[_lastFrame.Header.Version == SslProtocols.Ssl2 ? HandshakeTypeOffsetSsl2 : HandshakeTypeOffsetTls] == (byte)TlsHandshakeType.ClientHello; -#pragma warning restore CS0618 + _buffer.EncryptedReadOnlySpan[handshakeTypeOffset] == (byte)TlsHandshakeType.ClientHello; if (!isClientHello) { throw new AuthenticationException(SR.net_ssl_io_frame); diff --git a/src/libraries/System.Net.Security/tests/FunctionalTests/ServerAsyncAuthenticateTest.cs b/src/libraries/System.Net.Security/tests/FunctionalTests/ServerAsyncAuthenticateTest.cs index 949c0486e67f8e..ca7ebec187c299 100644 --- a/src/libraries/System.Net.Security/tests/FunctionalTests/ServerAsyncAuthenticateTest.cs +++ b/src/libraries/System.Net.Security/tests/FunctionalTests/ServerAsyncAuthenticateTest.cs @@ -369,6 +369,116 @@ public async Task ServerAsyncAuthenticate_InvalidHello_Throws(bool close) } } + public enum ServerCertificateSource + { + Direct, + SelectionCallback, + OptionsCallback + } + + public static IEnumerable EmptyHandshakeRecordData() + { + foreach (ServerCertificateSource certificateSource in Enum.GetValues()) + { + foreach (bool useAsync in new[] { false, true }) + { + if (!useAsync && certificateSource == ServerCertificateSource.OptionsCallback) + { + continue; + } + + yield return new object[] { certificateSource, useAsync, int.MaxValue, false }; + yield return new object[] { certificateSource, useAsync, 1, false }; + yield return new object[] { certificateSource, useAsync, int.MaxValue, true }; + } + } + } + + [Theory] + [MemberData(nameof(EmptyHandshakeRecordData))] + public Task ServerAuthenticate_EmptyHandshakeRecord_ThrowsIOException( + ServerCertificateSource certificateSource, bool useAsync, int maxReadSize, bool trailingData) => + AuthenticateEmptyHandshakeRecord(_serverCertificate, certificateSource, useAsync, maxReadSize, trailingData); + + internal static async Task AuthenticateEmptyHandshakeRecord( + X509Certificate2 certificate, ServerCertificateSource certificateSource, bool useAsync, int maxReadSize, bool trailingData) + { + byte[] record = trailingData + ? [0x16, 0x03, 0x01, 0x00, 0x00, 0x01] + : [0x16, 0x03, 0x01, 0x00, 0x00]; + using var input = new MemoryStream(record); + using var transport = new DelegateDelegatingStream(input) + { + ReadSpanFunc = Read, + ReadAsyncMemoryFunc = (buffer, _) => new ValueTask(Read(buffer.Span)) + }; + using var server = new SslStream(transport); + bool callbackInvoked = false; + var options = new SslServerAuthenticationOptions + { + // Exercise managed framing rather than Apple's Network Framework handshake. + EnabledSslProtocols = SslProtocols.Tls12 + }; + if (certificateSource == ServerCertificateSource.SelectionCallback) + { + options.ServerCertificateSelectionCallback = (_, _) => + { + callbackInvoked = true; + return certificate; + }; + } + else + { + options.ServerCertificate = certificate; + } + + if (certificateSource == ServerCertificateSource.OptionsCallback) + { + await Assert.ThrowsAsync(() => server.AuthenticateAsServerAsync((_, _, _, _) => + { + callbackInvoked = true; + return new ValueTask(options); + }, null)); + } + else if (useAsync) + { + await Assert.ThrowsAsync(() => server.AuthenticateAsServerAsync(options)); + } + else + { + Assert.Throws(() => server.AuthenticateAsServer(options)); + } + + Assert.False(callbackInvoked); + Assert.False(server.IsAuthenticated); + + int Read(Span buffer) + { + // Reject the record without an additional read that could mask the failure with EOF. + Assert.True(buffer.IsEmpty || input.Position < input.Length); + return input.Read(buffer.Slice(0, Math.Min(buffer.Length, maxReadSize))); + } + } + + [Fact] + public async Task ServerAsyncAuthenticate_EmptyHandshakeRecordWithoutEof_ThrowsIOException() + { + (Stream client, Stream server) = TestHelper.GetConnectedStreams(); + using (client) + using (var ssl = new SslStream(server)) + { + await client.WriteAsync(new byte[] { 0x16, 0x03, 0x01, 0x00, 0x00 }); + await Assert.ThrowsAsync(() => + ssl.AuthenticateAsServerAsync(new SslServerAuthenticationOptions + { + ServerCertificate = _serverCertificate, + EnabledSslProtocols = SslProtocols.Tls12 + }) + .WaitAsync(TestConfiguration.PassingTestTimeout)); + Assert.False(ssl.IsAuthenticated); + } + } + public static IEnumerable ProtocolMismatchData() { var supportedProtocols = new SslProtocolSupport.SupportedSslProtocolsTestData(); diff --git a/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamRemoteExecutorTests.cs b/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamRemoteExecutorTests.cs index dc20efda5d47ae..3b548096a32336 100644 --- a/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamRemoteExecutorTests.cs +++ b/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamRemoteExecutorTests.cs @@ -82,6 +82,24 @@ await TestConfiguration.WhenAllOrAnyFailedWithTimeout( }, appContextValue.ToString(), expectLegacyPath.ToString(), new RemoteInvokeOptions { StartInfo = psi }).DisposeAsync(); } + [ConditionalTheory(typeof(RemoteExecutor), nameof(RemoteExecutor.IsSupported))] + [PlatformSpecific(TestPlatforms.Windows | TestPlatforms.Linux | TestPlatforms.FreeBSD)] + [InlineData(false)] + [InlineData(true)] + public async Task ServerAuthenticate_EmptyHandshakeRecord_ThrowsIOException(bool useLegacyHandshake) + { + await RemoteExecutor.Invoke(async useLegacyHandshakeValue => + { + AppContext.SetSwitch("System.Net.Security.UseLegacySslStreamHandshake", bool.Parse(useLegacyHandshakeValue)); + using X509Certificate2 certificate = Configuration.Certificates.GetServerCertificate(); + foreach (object[] data in ServerAsyncAuthenticateTest.EmptyHandshakeRecordData()) + { + await ServerAsyncAuthenticateTest.AuthenticateEmptyHandshakeRecord( + certificate, (ServerAsyncAuthenticateTest.ServerCertificateSource)data[0], (bool)data[1], (int)data[2], (bool)data[3]); + } + }, useLegacyHandshake.ToString()).DisposeAsync(); + } + [ConditionalTheory(typeof(RemoteExecutor), nameof(RemoteExecutor.IsSupported))] [PlatformSpecific(TestPlatforms.Linux)] // SSLKEYLOGFILE is only supported on Linux for SslStream [InlineData(true)]