diff --git a/src/SharpCompress/IO/RewindableStream.Async.cs b/src/SharpCompress/IO/RewindableStream.Async.cs new file mode 100644 index 00000000..300874c5 --- /dev/null +++ b/src/SharpCompress/IO/RewindableStream.Async.cs @@ -0,0 +1,108 @@ +using System; +using System.IO; +using System.Threading; +using System.Threading.Tasks; + +namespace SharpCompress.IO; + +internal partial class RewindableStream +{ + public override async Task ReadAsync( + byte[] buffer, + int offset, + int count, + CancellationToken cancellationToken + ) + { + //don't actually read if we don't really want to read anything + //currently a network stream bug on Windows for .NET Core + if (count == 0) + { + return 0; + } + int read; + if (_isRewound && _bufferStream.Position != _bufferStream.Length) + { + read = await _bufferStream + .ReadAsync(buffer, offset, count, cancellationToken) + .ConfigureAwait(false); + if (read < count) + { + int tempRead = await stream + .ReadAsync(buffer, offset + read, count - read, cancellationToken) + .ConfigureAwait(false); + if (IsRecording) + { + await _bufferStream + .WriteAsync(buffer, offset + read, tempRead, cancellationToken) + .ConfigureAwait(false); + } + read += tempRead; + } + if (_bufferStream.Position == _bufferStream.Length && !IsRecording) + { + _isRewound = false; + _bufferStream.SetLength(0); + } + return read; + } + + read = await stream + .ReadAsync(buffer, offset, count, cancellationToken) + .ConfigureAwait(false); + if (IsRecording) + { + await _bufferStream + .WriteAsync(buffer, offset, read, cancellationToken) + .ConfigureAwait(false); + } + return read; + } + +#if !LEGACY_DOTNET + public override async ValueTask ReadAsync( + Memory buffer, + CancellationToken cancellationToken = default + ) + { + //don't actually read if we don't really want to read anything + //currently a network stream bug on Windows for .NET Core + if (buffer.Length == 0) + { + return 0; + } + int read; + if (_isRewound && _bufferStream.Position != _bufferStream.Length) + { + var bufferSpan = buffer.Span; + read = _bufferStream.Read(bufferSpan); + if (read < bufferSpan.Length) + { + int tempRead = await stream + .ReadAsync(buffer.Slice(read), cancellationToken) + .ConfigureAwait(false); + if (IsRecording) + { + await _bufferStream + .WriteAsync(buffer.Slice(read, tempRead), cancellationToken) + .ConfigureAwait(false); + } + read += tempRead; + } + if (_bufferStream.Position == _bufferStream.Length && !IsRecording) + { + _isRewound = false; + _bufferStream.SetLength(0); + } + return read; + } + + read = await stream.ReadAsync(buffer, cancellationToken).ConfigureAwait(false); + if (IsRecording) + { + await _bufferStream.WriteAsync(buffer.Slice(0, read), cancellationToken).ConfigureAwait(false); + } + return read; + } +#endif +} diff --git a/src/SharpCompress/IO/RewindableStream.cs b/src/SharpCompress/IO/RewindableStream.cs index d265fe4d..86f445e9 100644 --- a/src/SharpCompress/IO/RewindableStream.cs +++ b/src/SharpCompress/IO/RewindableStream.cs @@ -1,29 +1,23 @@ -using System; +using System; using System.IO; namespace SharpCompress.IO { - internal partial class RewindableStream : Stream + internal partial class RewindableStream(Stream stream) : Stream { - private readonly Stream stream; - private MemoryStream bufferStream = new MemoryStream(); - private bool isRewound; - private bool isDisposed; - - public RewindableStream(Stream stream) - { - this.stream = stream; - } + private MemoryStream _bufferStream = new MemoryStream(); + private bool _isRewound; + private bool _isDisposed; internal bool IsRecording { get; private set; } protected override void Dispose(bool disposing) { - if (isDisposed) + if (_isDisposed) { return; } - isDisposed = true; + _isDisposed = true; base.Dispose(disposing); if (disposing) { @@ -33,42 +27,41 @@ namespace SharpCompress.IO public void Rewind(bool stopRecording = false) { - isRewound = true; + _isRewound = true; IsRecording = !stopRecording; - bufferStream.Position = 0; + _bufferStream.Position = 0; } public void Rewind(MemoryStream buffer) { - if (bufferStream.Position >= buffer.Length) + if (_bufferStream.Position >= buffer.Length) { - bufferStream.Position -= buffer.Length; + _bufferStream.Position -= buffer.Length; } else { - - bufferStream.TransferTo(buffer, buffer.Length - bufferStream.Position); + _bufferStream.TransferTo(buffer, buffer.Length - _bufferStream.Position); //create new memorystream to allow proper resizing as memorystream could be a user provided buffer //https://github.com/adamhathcock/sharpcompress/issues/306 - bufferStream = new MemoryStream(); + _bufferStream = new MemoryStream(); buffer.Position = 0; - buffer.TransferTo(bufferStream, buffer.Length); - bufferStream.Position = 0; + buffer.TransferTo(_bufferStream, buffer.Length); + _bufferStream.Position = 0; } - isRewound = true; + _isRewound = true; } public void StartRecording() { //if (isRewound && bufferStream.Position != 0) // throw new System.NotImplementedException(); - if (bufferStream.Position != 0) + if (_bufferStream.Position != 0) { - byte[] data = bufferStream.ToArray(); - long position = bufferStream.Position; - bufferStream.SetLength(0); - bufferStream.Write(data, (int)position, data.Length - (int)position); - bufferStream.Position = 0; + byte[] data = _bufferStream.ToArray(); + long position = _bufferStream.Position; + _bufferStream.SetLength(0); + _bufferStream.Write(data, (int)position, data.Length - (int)position); + _bufferStream.Position = 0; } IsRecording = true; } @@ -79,31 +72,28 @@ namespace SharpCompress.IO public override bool CanWrite => false; - public override void Flush() - { - throw new NotSupportedException(); - } + public override void Flush() => throw new NotSupportedException(); public override long Length => throw new NotSupportedException(); public override long Position { - get { return stream.Position + bufferStream.Position - bufferStream.Length; } + get => stream.Position + _bufferStream.Position - _bufferStream.Length; set { - if (!isRewound) + if (!_isRewound) { stream.Position = value; } - else if (value < stream.Position - bufferStream.Length || value >= stream.Position) + else if (value < stream.Position - _bufferStream.Length || value >= stream.Position) { stream.Position = value; - isRewound = false; - bufferStream.SetLength(0); + _isRewound = false; + _bufferStream.SetLength(0); } else { - bufferStream.Position = value - stream.Position + bufferStream.Length; + _bufferStream.Position = value - stream.Position + _bufferStream.Length; } } } @@ -117,22 +107,22 @@ namespace SharpCompress.IO return 0; } int read; - if (isRewound && bufferStream.Position != bufferStream.Length) + if (_isRewound && _bufferStream.Position != _bufferStream.Length) { - read = bufferStream.Read(buffer, offset, count); + read = _bufferStream.Read(buffer, offset, count); if (read < count) { int tempRead = stream.Read(buffer, offset + read, count - read); if (IsRecording) { - bufferStream.Write(buffer, offset + read, tempRead); + _bufferStream.Write(buffer, offset + read, tempRead); } read += tempRead; } - if (bufferStream.Position == bufferStream.Length && !IsRecording) + if (_bufferStream.Position == _bufferStream.Length && !IsRecording) { - isRewound = false; - bufferStream.SetLength(0); + _isRewound = false; + _bufferStream.SetLength(0); } return read; } @@ -140,24 +130,54 @@ namespace SharpCompress.IO read = stream.Read(buffer, offset, count); if (IsRecording) { - bufferStream.Write(buffer, offset, read); + _bufferStream.Write(buffer, offset, read); } return read; } - public override long Seek(long offset, SeekOrigin origin) +#if !LEGACY_DOTNET + public override int Read(Span buffer) { - throw new NotSupportedException(); - } + //don't actually read if we don't really want to read anything + //currently a network stream bug on Windows for .NET Core + if (buffer.Length == 0) + { + return 0; + } + int read; + if (_isRewound && _bufferStream.Position != _bufferStream.Length) + { + read = _bufferStream.Read(buffer); + if (read < buffer.Length) + { + int tempRead = stream.Read(buffer.Slice(read)); + if (IsRecording) + { + _bufferStream.Write(buffer.Slice(read, tempRead)); + } + read += tempRead; + } + if (_bufferStream.Position == _bufferStream.Length && !IsRecording) + { + _isRewound = false; + _bufferStream.SetLength(0); + } + return read; + } - public override void SetLength(long value) - { - throw new NotSupportedException(); + read = stream.Read(buffer); + if (IsRecording) + { + _bufferStream.Write(buffer.Slice(0, read)); + } + return read; } +#endif - public override void Write(byte[] buffer, int offset, int count) - { - throw new NotSupportedException(); - } + public override long Seek(long offset, SeekOrigin origin) => throw new NotSupportedException(); + + public override void SetLength(long value) => throw new NotSupportedException(); + + public override void Write(byte[] buffer, int offset, int count) => throw new NotSupportedException(); } } diff --git a/tests/SharpCompress.Test/Streams/RewindableStreamAsyncTest.cs b/tests/SharpCompress.Test/Streams/RewindableStreamAsyncTest.cs index 455f007f..692b5b01 100644 --- a/tests/SharpCompress.Test/Streams/RewindableStreamAsyncTest.cs +++ b/tests/SharpCompress.Test/Streams/RewindableStreamAsyncTest.cs @@ -1,4 +1,7 @@ +using System; +using System.Collections.Generic; using System.IO; +using System.Threading; using System.Threading.Tasks; using SharpCompress.IO; using Xunit; @@ -21,15 +24,14 @@ public class RewindableStreamAsyncTest bw.Write(7); bw.Flush(); ms.Position = 0; - var stream = new SharpCompressStream(ms, bufferSize: 0x10000); + var stream = new RewindableStream(ms); + stream.StartRecording(); Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(2, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(3, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(4, await ReadInt32Async(stream).ConfigureAwait(false)); - - ((IStreamStack)stream).StackSeek(0); - long pos = stream.Position; + stream.Rewind(true); Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(2, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(3, await ReadInt32Async(stream).ConfigureAwait(false)); @@ -37,12 +39,6 @@ public class RewindableStreamAsyncTest Assert.Equal(5, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(6, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(7, await ReadInt32Async(stream).ConfigureAwait(false)); - - ((IStreamStack)stream).StackSeek(pos); - Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false)); - Assert.Equal(2, await ReadInt32Async(stream).ConfigureAwait(false)); - Assert.Equal(3, await ReadInt32Async(stream).ConfigureAwait(false)); - Assert.Equal(4, await ReadInt32Async(stream).ConfigureAwait(false)); } [Fact] @@ -59,23 +55,16 @@ public class RewindableStreamAsyncTest bw.Write(7); bw.Flush(); ms.Position = 0; - var stream = new SharpCompressStream(ms, bufferSize: 0x10000); + var stream = new RewindableStream(ms); + stream.StartRecording(); Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(2, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(3, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(4, await ReadInt32Async(stream).ConfigureAwait(false)); - ((IStreamStack)stream).StackSeek(0); - + stream.Rewind(true); Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(2, await ReadInt32Async(stream).ConfigureAwait(false)); - long pos = stream.Position; - - Assert.Equal(3, await ReadInt32Async(stream).ConfigureAwait(false)); - Assert.Equal(4, await ReadInt32Async(stream).ConfigureAwait(false)); - Assert.Equal(5, await ReadInt32Async(stream).ConfigureAwait(false)); - ((IStreamStack)stream).StackSeek(pos); - Assert.Equal(3, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(4, await ReadInt32Async(stream).ConfigureAwait(false)); Assert.Equal(5, await ReadInt32Async(stream).ConfigureAwait(false)); @@ -83,6 +72,254 @@ public class RewindableStreamAsyncTest Assert.Equal(7, await ReadInt32Async(stream).ConfigureAwait(false)); } + [Fact] + public async ValueTask TestRecordingAsync() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + bw.Write(1); + bw.Write(2); + bw.Write(3); + bw.Write(4); + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + stream.StartRecording(); + + Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false)); + Assert.Equal(2, await ReadInt32Async(stream).ConfigureAwait(false)); + stream.Rewind(false); + Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false)); + Assert.Equal(2, await ReadInt32Async(stream).ConfigureAwait(false)); + Assert.Equal(3, await ReadInt32Async(stream).ConfigureAwait(false)); + Assert.Equal(4, await ReadInt32Async(stream).ConfigureAwait(false)); + } + + [Fact] + public async ValueTask TestAsyncProducesSameResultAsSync() + { + var testData = new byte[100 * 4]; + for (int i = 0; i < 100; i++) + { + var bytes = BitConverter.GetBytes(i); + Array.Copy(bytes, 0, testData, i * 4, 4); + } + + byte[] syncResult; + byte[] asyncResult; + + var ms1 = new MemoryStream(testData); + using (var stream = new RewindableStream(ms1)) + { + syncResult = ReadAllSync(stream); + } + + var ms2 = new MemoryStream(testData); + using (var stream = new RewindableStream(ms2)) + { + asyncResult = await ReadAllAsync(stream).ConfigureAwait(false); + } + + Assert.Equal(syncResult, asyncResult); + } + + [Fact] + public async ValueTask TestAsyncWithRewind() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + for (int i = 0; i < 50; i++) + { + bw.Write(i); + } + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + stream.StartRecording(); + + var buffer = new byte[8]; + await stream.ReadAsync(buffer, 0, buffer.Length).ConfigureAwait(false); + Assert.Equal(0, BitConverter.ToInt32(buffer, 0)); + Assert.Equal(1, BitConverter.ToInt32(buffer, 4)); + + stream.Rewind(false); + await stream.ReadAsync(buffer, 0, buffer.Length).ConfigureAwait(false); + Assert.Equal(0, BitConverter.ToInt32(buffer, 0)); + Assert.Equal(1, BitConverter.ToInt32(buffer, 4)); + + await stream.ReadAsync(buffer, 0, buffer.Length).ConfigureAwait(false); + Assert.Equal(2, BitConverter.ToInt32(buffer, 0)); + Assert.Equal(3, BitConverter.ToInt32(buffer, 4)); + } + + [Fact] + public async ValueTask TestAsyncCancellationSupport() + { + var ms = new MemoryStream(new byte[10000]); + ms.Position = 0; + var stream = new RewindableStream(ms); + + var cts = new CancellationTokenSource(); + var buffer = new byte[4096]; + + // Just verify that cancellation token can be passed without throwing + int bytesRead = await stream + .ReadAsync(buffer, 0, buffer.Length, cts.Token) + .ConfigureAwait(false); + Assert.Equal(buffer.Length, bytesRead); + } + + [Fact] + public async ValueTask TestAsyncEmptyBuffer() + { + var ms = new MemoryStream(); + ms.Position = 0; + var stream = new RewindableStream(ms); + + var buffer = new byte[0]; + int bytesRead = await stream.ReadAsync(buffer, 0, 0).ConfigureAwait(false); + Assert.Equal(0, bytesRead); + } + + [Fact] + public async ValueTask TestAsyncMultipleReads() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + for (int i = 0; i < 50; i++) + { + bw.Write(i); + } + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + + var totalData = new byte[50 * 4]; + var buffer = new byte[8]; + int offset = 0; + int bytesRead; + + while ( + (bytesRead = await stream.ReadAsync(buffer, 0, buffer.Length).ConfigureAwait(false)) > 0 + ) + { + Array.Copy(buffer, 0, totalData, offset, bytesRead); + offset += bytesRead; + } + + Assert.Equal(50 * 4, offset); + Assert.Equal(0, BitConverter.ToInt32(totalData, 0)); + Assert.Equal(49, BitConverter.ToInt32(totalData, 49 * 4)); + } + + [Fact] + public async ValueTask TestAsyncReturnsZeroAtEndOfStream() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + bw.Write(1); + bw.Write(2); + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + + var buffer = new byte[4096]; + + int bytesRead; + while ( + (bytesRead = await stream.ReadAsync(buffer, 0, buffer.Length).ConfigureAwait(false)) > 0 + ) { } + + Assert.Equal(0, bytesRead); + + bytesRead = await stream.ReadAsync(buffer, 0, buffer.Length).ConfigureAwait(false); + Assert.Equal(0, bytesRead); + } + + [Fact] + public async ValueTask TestAsyncPosition() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + for (int i = 0; i < 10; i++) + { + bw.Write(i); + } + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + Assert.Equal(0, stream.Position); + + var buffer = new byte[4]; + await stream.ReadAsync(buffer, 0, 4).ConfigureAwait(false); + Assert.Equal(4, stream.Position); + + stream.StartRecording(); + await stream.ReadAsync(buffer, 0, 4).ConfigureAwait(false); + Assert.Equal(8, stream.Position); + + stream.Rewind(); + Assert.Equal(4, stream.Position); + } + +#if !LEGACY_DOTNET + [Fact] + public async ValueTask TestAsyncMemoryCancellationSupport() + { + var ms = new MemoryStream(new byte[10000]); + ms.Position = 0; + var stream = new RewindableStream(ms); + + var cts = new CancellationTokenSource(); + var buffer = new byte[4096]; + + // Just verify that cancellation token can be passed without throwing + int bytesRead = await stream.ReadAsync(buffer.AsMemory(), cts.Token).ConfigureAwait(false); + Assert.Equal(buffer.Length, bytesRead); + } + + [Fact] + public async ValueTask TestAsyncMemoryEmptyBuffer() + { + var ms = new MemoryStream(); + ms.Position = 0; + var stream = new RewindableStream(ms); + + var buffer = Memory.Empty; + int bytesRead = await stream.ReadAsync(buffer).ConfigureAwait(false); + Assert.Equal(0, bytesRead); + } + + [Fact] + public async ValueTask TestAsyncMemoryMultipleReads() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + for (int i = 0; i < 50; i++) + { + bw.Write(i); + } + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + + var totalData = new byte[50 * 4]; + var buffer = new byte[8]; + int offset = 0; + int bytesRead; + + while ((bytesRead = await stream.ReadAsync(buffer.AsMemory()).ConfigureAwait(false)) > 0) + { + Array.Copy(buffer, 0, totalData, offset, bytesRead); + offset += bytesRead; + } + + Assert.Equal(50 * 4, offset); + Assert.Equal(0, BitConverter.ToInt32(totalData, 0)); + Assert.Equal(49, BitConverter.ToInt32(totalData, 49 * 4)); + } +#endif + private static async Task ReadInt32Async(Stream stream) { var buffer = new byte[4]; @@ -93,4 +330,72 @@ public class RewindableStreamAsyncTest } return buffer[0] | (buffer[1] << 8) | (buffer[2] << 16) | (buffer[3] << 24); } + +#if !LEGACY_DOTNET + private static async ValueTask ReadInt32AsyncMemory(Stream stream) + { + var buffer = new byte[4]; + var bytesRead = await stream.ReadAsync(buffer.AsMemory()).ConfigureAwait(false); + if (bytesRead != 4) + { + throw new EndOfStreamException(); + } + return buffer[0] | (buffer[1] << 8) | (buffer[2] << 16) | (buffer[3] << 24); + } +#endif + + private static byte[] ReadAllSync(RewindableStream stream) + { + var result = new List(); + var buffer = new byte[4096]; + int bytesRead; + + while ((bytesRead = stream.Read(buffer, 0, buffer.Length)) > 0) + { + for (int i = 0; i < bytesRead; i++) + { + result.Add(buffer[i]); + } + } + + return result.ToArray(); + } + + private static async Task ReadAllAsync(RewindableStream stream) + { + var result = new List(); + var buffer = new byte[4096]; + int bytesRead; + + while ( + (bytesRead = await stream.ReadAsync(buffer, 0, buffer.Length).ConfigureAwait(false)) > 0 + ) + { + for (int i = 0; i < bytesRead; i++) + { + result.Add(buffer[i]); + } + } + + return result.ToArray(); + } + +#if !LEGACY_DOTNET + private static async ValueTask ReadAllAsyncMemory(RewindableStream stream) + { + var result = new List(); + var buffer = new byte[4096]; + int bytesRead; + + while ((bytesRead = await stream.ReadAsync(buffer.AsMemory()).ConfigureAwait(false)) > 0) + { + for (int i = 0; i < bytesRead; i++) + { + result.Add(buffer[i]); + } + } + + return result.ToArray(); + } +#endif } diff --git a/tests/SharpCompress.Test/Streams/RewindableStreamTest.cs b/tests/SharpCompress.Test/Streams/RewindableStreamTest.cs index 37580df9..bf6a74a2 100644 --- a/tests/SharpCompress.Test/Streams/RewindableStreamTest.cs +++ b/tests/SharpCompress.Test/Streams/RewindableStreamTest.cs @@ -1,3 +1,4 @@ +using System; using System.IO; using SharpCompress.IO; using Xunit; @@ -20,17 +21,14 @@ public class RewindableStreamTest bw.Write(7); bw.Flush(); ms.Position = 0; - var stream = new SharpCompressStream(ms, bufferSize: 0x10000); - //stream.StartRecording(); + var stream = new RewindableStream(ms); + stream.StartRecording(); var br = new BinaryReader(stream); Assert.Equal(1, br.ReadInt32()); Assert.Equal(2, br.ReadInt32()); Assert.Equal(3, br.ReadInt32()); Assert.Equal(4, br.ReadInt32()); - //stream.Rewind(true); - ((IStreamStack)stream).StackSeek(0); - //stream.StartRecording(); - long pos = stream.Position; + stream.Rewind(true); Assert.Equal(1, br.ReadInt32()); Assert.Equal(2, br.ReadInt32()); Assert.Equal(3, br.ReadInt32()); @@ -38,13 +36,6 @@ public class RewindableStreamTest Assert.Equal(5, br.ReadInt32()); Assert.Equal(6, br.ReadInt32()); Assert.Equal(7, br.ReadInt32()); - //stream.Rewind(true); - //stream.StartRecording(); - ((IStreamStack)stream).StackSeek(pos); - Assert.Equal(1, br.ReadInt32()); - Assert.Equal(2, br.ReadInt32()); - Assert.Equal(3, br.ReadInt32()); - Assert.Equal(4, br.ReadInt32()); } [Fact] @@ -61,24 +52,16 @@ public class RewindableStreamTest bw.Write(7); bw.Flush(); ms.Position = 0; - var stream = new SharpCompressStream(ms, bufferSize: 0x10000); - //stream.StartRecording(); + var stream = new RewindableStream(ms); + stream.StartRecording(); var br = new BinaryReader(stream); Assert.Equal(1, br.ReadInt32()); Assert.Equal(2, br.ReadInt32()); Assert.Equal(3, br.ReadInt32()); Assert.Equal(4, br.ReadInt32()); - ((IStreamStack)stream).StackSeek(0); - //stream.Rewind(true); + stream.Rewind(true); Assert.Equal(1, br.ReadInt32()); Assert.Equal(2, br.ReadInt32()); - long pos = stream.Position; - //stream.StartRecording(); - Assert.Equal(3, br.ReadInt32()); - Assert.Equal(4, br.ReadInt32()); - Assert.Equal(5, br.ReadInt32()); - ((IStreamStack)stream).StackSeek(pos); - //stream.Rewind(true); Assert.Equal(3, br.ReadInt32()); Assert.Equal(4, br.ReadInt32()); Assert.Equal(5, br.ReadInt32()); @@ -87,20 +70,128 @@ public class RewindableStreamTest } [Fact] - public void TestSmallBuffer() + public void TestRecording() { var ms = new MemoryStream(); - var testData = new byte[100]; - for (byte i = 0; i < 100; i++) - { - testData[i] = i; - } - ms.Write(testData); + var bw = new BinaryWriter(ms); + bw.Write(1); + bw.Write(2); + bw.Write(3); + bw.Write(4); + bw.Flush(); ms.Position = 0; - using var stream = new SharpCompressStream(ms, bufferSize: 64); + var stream = new RewindableStream(ms); + stream.StartRecording(); var br = new BinaryReader(stream); - stream.StackSeek(100); - stream.StackSeek(10); - Assert.Equal(10, br.ReadByte()); + Assert.Equal(1, br.ReadInt32()); + Assert.Equal(2, br.ReadInt32()); + stream.Rewind(false); + Assert.Equal(1, br.ReadInt32()); + Assert.Equal(2, br.ReadInt32()); + Assert.Equal(3, br.ReadInt32()); + Assert.Equal(4, br.ReadInt32()); + } + + [Fact] + public void TestRewindWithPartialBuffer() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + bw.Write(1); + bw.Write(2); + bw.Write(3); + bw.Write(4); + bw.Write(5); + bw.Write(6); + bw.Write(7); + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + stream.StartRecording(); + var br = new BinaryReader(stream); + Assert.Equal(1, br.ReadInt32()); + Assert.Equal(2, br.ReadInt32()); + Assert.Equal(3, br.ReadInt32()); + Assert.Equal(4, br.ReadInt32()); + Assert.Equal(5, br.ReadInt32()); + Assert.Equal(6, br.ReadInt32()); + + // Create a buffer with the last 2 ints (12-16 bytes) + var externalBuffer = new MemoryStream(); + externalBuffer.Write(BitConverter.GetBytes(3), 0, 4); + externalBuffer.Write(BitConverter.GetBytes(4), 0, 4); + externalBuffer.Write(BitConverter.GetBytes(5), 0, 4); + externalBuffer.Write(BitConverter.GetBytes(6), 0, 4); + externalBuffer.Position = 0; + + // Rewind by 12 bytes (3 ints worth) + stream.Rewind(externalBuffer); + Assert.Equal(3, br.ReadInt32()); + Assert.Equal(4, br.ReadInt32()); + Assert.Equal(5, br.ReadInt32()); + Assert.Equal(6, br.ReadInt32()); + Assert.Equal(7, br.ReadInt32()); + } + + [Fact] + public void TestPosition() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + for (int i = 0; i < 10; i++) + { + bw.Write(i); + } + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + Assert.Equal(0, stream.Position); + + var buffer = new byte[4]; + stream.Read(buffer, 0, 4); + Assert.Equal(4, stream.Position); + + stream.StartRecording(); + stream.Read(buffer, 0, 4); + Assert.Equal(8, stream.Position); + + stream.Rewind(); + Assert.Equal(4, stream.Position); + } + + [Fact] + public void TestPositionSeek() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + for (int i = 0; i < 10; i++) + { + bw.Write(i); + } + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + var br = new BinaryReader(stream); + + Assert.Equal(0, br.ReadInt32()); + Assert.Equal(1, br.ReadInt32()); + Assert.Equal(2, br.ReadInt32()); + + stream.Position = 4; + Assert.Equal(1, br.ReadInt32()); + } + + [Fact] + public void TestDispose() + { + var ms = new MemoryStream(); + var bw = new BinaryWriter(ms); + bw.Write(1); + bw.Write(2); + bw.Flush(); + ms.Position = 0; + var stream = new RewindableStream(ms); + stream.Dispose(); + Assert.Throws(() => stream.Read(new byte[4], 0, 4)); } }