From 9fd0c2b83e0792acfe36d7c13a14bf479277a698 Mon Sep 17 00:00:00 2001 From: Stephen Toub Date: Thu, 28 May 2020 10:55:33 -0400 Subject: [PATCH 1/2] Add SslStream test for zero-byte reads We have higher-level tests for libraries like WebSockets that validate the behavior of zero-byte reads, but we don't currently have one in System.Net.Security's tests. Adding one. --- .../SslStreamStreamToStreamTest.cs | 32 +++++++++++++++++++ 1 file changed, 32 insertions(+) diff --git a/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamStreamToStreamTest.cs b/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamStreamToStreamTest.cs index 9fb2ac9fefbd6b..6d8f106dac6e94 100644 --- a/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamStreamToStreamTest.cs +++ b/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamStreamToStreamTest.cs @@ -259,6 +259,38 @@ public async Task SslStream_StreamToStream_Successive_ClientWrite_WithZeroBytes_ } } + [Fact] + public async Task SslStream_StreamToStream_ZeroByteRead_SucceedsWhenDataAvailable() + { + using var listener = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp); + using var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp); + listener.Bind(new IPEndPoint(IPAddress.Loopback, 0)); + listener.Listen(1); + client.Connect(listener.LocalEndPoint); + using Socket server = listener.Accept(); + + using var clientSslStream = new SslStream(new NetworkStream(client, ownsSocket: true), leaveInnerStreamOpen: false, AllowAnyServerCertificate); + using var serverSslStream = new SslStream(new NetworkStream(server, ownsSocket: true)); + await DoHandshake(clientSslStream, serverSslStream); + + for (int iter = 0; iter < 2; iter++) + { + ValueTask zeroByteRead = clientSslStream.ReadAsync(Memory.Empty); + Assert.False(zeroByteRead.IsCompleted); + + await serverSslStream.WriteAsync(Encoding.UTF8.GetBytes("hello")); + Assert.Equal(0, await zeroByteRead); + + var readBytes = new byte[5]; + int count = 0; + while (count < readBytes.Length) + { + count += await clientSslStream.ReadAsync(readBytes.AsMemory(count)); + } + Assert.Equal("hello", Encoding.UTF8.GetString(readBytes)); + } + } + [Theory] [InlineData(false)] [InlineData(true)] From 502e169d9879ba36c81c8a744ba190f84dfdedc0 Mon Sep 17 00:00:00 2001 From: Stephen Toub Date: Fri, 29 May 2020 17:17:04 -0400 Subject: [PATCH 2/2] Address PR feedback --- .../SslStreamStreamToStreamTest.cs | 16 ++++++---------- 1 file changed, 6 insertions(+), 10 deletions(-) diff --git a/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamStreamToStreamTest.cs b/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamStreamToStreamTest.cs index 6d8f106dac6e94..d04f4b56c38438 100644 --- a/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamStreamToStreamTest.cs +++ b/src/libraries/System.Net.Security/tests/FunctionalTests/SslStreamStreamToStreamTest.cs @@ -262,15 +262,9 @@ public async Task SslStream_StreamToStream_Successive_ClientWrite_WithZeroBytes_ [Fact] public async Task SslStream_StreamToStream_ZeroByteRead_SucceedsWhenDataAvailable() { - using var listener = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp); - using var client = new Socket(AddressFamily.InterNetwork, SocketType.Stream, ProtocolType.Tcp); - listener.Bind(new IPEndPoint(IPAddress.Loopback, 0)); - listener.Listen(1); - client.Connect(listener.LocalEndPoint); - using Socket server = listener.Accept(); - - using var clientSslStream = new SslStream(new NetworkStream(client, ownsSocket: true), leaveInnerStreamOpen: false, AllowAnyServerCertificate); - using var serverSslStream = new SslStream(new NetworkStream(server, ownsSocket: true)); + (NetworkStream clientStream, NetworkStream serverStream) = TestHelper.GetConnectedTcpStreams(); + using var clientSslStream = new SslStream(clientStream, leaveInnerStreamOpen: false, AllowAnyServerCertificate); + using var serverSslStream = new SslStream(serverStream); await DoHandshake(clientSslStream, serverSslStream); for (int iter = 0; iter < 2; iter++) @@ -285,7 +279,9 @@ public async Task SslStream_StreamToStream_ZeroByteRead_SucceedsWhenDataAvailabl int count = 0; while (count < readBytes.Length) { - count += await clientSslStream.ReadAsync(readBytes.AsMemory(count)); + int n = await clientSslStream.ReadAsync(readBytes.AsMemory(count)); + Assert.InRange(n, 1, readBytes.Length - count); + count += n; } Assert.Equal("hello", Encoding.UTF8.GetString(readBytes)); }