making RewindableStream more proper

This commit is contained in:
Adam Hathcock 2026-01-28 16:50:35 +00:00
parent c770bc4788
commit 8dfd5349f0
4 changed files with 635 additions and 111 deletions

View file

@ -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<int> 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<int> ReadAsync(
Memory<byte> 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
}

View file

@ -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<byte> 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();
}
}

View file

@ -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<byte>.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<int> 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<int> 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<byte>();
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<byte[]> ReadAllAsync(RewindableStream stream)
{
var result = new List<byte>();
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<byte[]> ReadAllAsyncMemory(RewindableStream stream)
{
var result = new List<byte>();
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
}

View file

@ -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<ObjectDisposedException>(() => stream.Read(new byte[4], 0, 4));
}
}