more rework

This commit is contained in:
Adam Hathcock
2026-01-29 14:42:29 +00:00
parent be8841075a
commit 1ca914823f
16 changed files with 1059 additions and 155 deletions

View File

@@ -153,11 +153,6 @@ public partial class ZipArchive
var headerFactory = new StreamingZipHeaderFactory(password, new ArchiveEncoding(), null);
try
{
if (stream is not SharpCompressStream)
{
stream = new SharpCompressStream(stream, bufferSize: Constants.BufferSize);
}
var header = headerFactory
.ReadStreamHeader(stream)
.FirstOrDefault(x => x.ZipHeaderType != ZipHeaderType.Split);

View File

@@ -26,7 +26,7 @@ internal sealed partial class StreamingZipFilePart : ZipFilePart
);
if (LeaveStreamOpen)
{
return SharpCompressStream.Create(_decompressionStream, leaveOpen: true);
return new NonDisposingStream(_decompressionStream);
}
return _decompressionStream;
}

View File

@@ -60,7 +60,7 @@ internal sealed partial class StreamingZipHeaderFactory
private sealed class StreamHeaderAsyncEnumerator : IAsyncEnumerator<ZipHeader>, IDisposable
{
private readonly StreamingZipHeaderFactory _headerFactory;
private readonly SharpCompressStream _rewindableStream;
private readonly RewindableStream _rewindableStream;
private readonly AsyncBinaryReader _reader;
private readonly CancellationToken _cancellationToken;
private bool _completed;
@@ -72,7 +72,7 @@ internal sealed partial class StreamingZipHeaderFactory
)
{
_headerFactory = headerFactory;
_rewindableStream = EnsureSharpCompressStream(stream);
_rewindableStream = RewindableStream.EnsureSeekable(stream);
_reader = new AsyncBinaryReader(_rewindableStream, leaveOpen: true);
_cancellationToken = cancellationToken;
}
@@ -281,10 +281,11 @@ internal sealed partial class StreamingZipHeaderFactory
} // Check if zip is streaming ( Length is 0 and is declared in PostDataDescriptor )
else if (localHeader.Flags.HasFlag(HeaderFlags.UsePostDataDescriptor))
{
_rewindableStream.StartRecording();
var nextHeaderBytes = await _reader
.ReadUInt32Async(_cancellationToken)
.ConfigureAwait(false);
((IStreamStack)_rewindableStream).Rewind(sizeof(uint));
_rewindableStream.Rewind(true);
// Check if next data is PostDataDescriptor, streamed file with 0 length
header.HasData = !IsHeader(nextHeaderBytes);
@@ -313,29 +314,5 @@ internal sealed partial class StreamingZipHeaderFactory
{
_reader.Dispose();
}
/// <summary>
/// Ensures the stream is a <see cref="SharpCompressStream"/> so header parsing can use rewind/buffer helpers.
/// </summary>
private static SharpCompressStream EnsureSharpCompressStream(Stream stream)
{
if (stream is SharpCompressStream sharpCompressStream)
{
return sharpCompressStream;
}
// Ensure the stream is already a SharpCompressStream so the buffer/size is set.
// The original code wrapped this with RewindableStream; use SharpCompressStream so we can get the buffer size.
if (stream is SourceStream src)
{
return new SharpCompressStream(
stream,
src.ReaderOptions.LeaveStreamOpen,
bufferSize: src.ReaderOptions.BufferSize
);
}
throw new ArgumentException("Stream must be a SharpCompressStream", nameof(stream));
}
}
}

View File

@@ -23,8 +23,7 @@ internal sealed partial class StreamingZipHeaderFactory : ZipHeaderFactory
internal IEnumerable<ZipHeader> ReadStreamHeader(Stream stream)
{
//the original code wrapped this with RewindableStream. Wrap with SharpCompressStream as we can get the buffer size
var rewindableStream = stream;
var rewindableStream = RewindableStream.EnsureSeekable(stream);
while (true)
{
var reader = new BinaryReader(rewindableStream);
@@ -175,8 +174,9 @@ internal sealed partial class StreamingZipHeaderFactory : ZipHeaderFactory
} // Check if zip is streaming ( Length is 0 and is declared in PostDataDescriptor )
else if (local_header.Flags.HasFlag(HeaderFlags.UsePostDataDescriptor))
{
rewindableStream.StartRecording();
var nextHeaderBytes = reader.ReadUInt32();
((IStreamStack)rewindableStream).Rewind(sizeof(uint));
rewindableStream.Rewind(true);
// Check if next data is PostDataDescriptor, streamed file with 0 length
header.HasData = !IsHeader(nextHeaderBytes);
@@ -189,25 +189,4 @@ internal sealed partial class StreamingZipHeaderFactory : ZipHeaderFactory
yield return header;
}
}
private static SharpCompressStream EnsureSharpCompressStream(Stream stream)
{
if (stream is SharpCompressStream sharpCompressStream)
{
return sharpCompressStream;
}
// Ensure the stream is already a SharpCompressStream so the buffer/size is set.
// The original code wrapped this with RewindableStream; use SharpCompressStream so we can get the buffer size.
if (stream is SourceStream src)
{
return new SharpCompressStream(
stream,
src.ReaderOptions.LeaveStreamOpen,
bufferSize: src.ReaderOptions.BufferSize
);
}
throw new ArgumentException("Stream must be a SharpCompressStream", nameof(stream));
}
}

View File

@@ -213,7 +213,7 @@ internal abstract partial class ZipFilePart : FilePart
) || Header.IsZip64
)
{
plainStream = SharpCompressStream.Create(plainStream, leaveOpen: true); //make sure AES doesn't close
plainStream = new NonDisposingStream(plainStream); //make sure AES doesn't close
}
else
{

View File

@@ -87,7 +87,7 @@ public abstract class Factory : IFactory
stream.Rewind();
if (IsArchive(stream, options.Password))
{
stream.StopRecording();
stream.Rewind(true);
reader = readerFactory.OpenReader(stream, options);
return true;
}

View File

@@ -44,14 +44,6 @@ public class ZipFactory
public override bool IsArchive(Stream stream, string? password = null)
{
var startPosition = stream.CanSeek ? stream.Position : -1;
// probe for single volume zip
if (stream is not SharpCompressStream) // wrap to provide buffer bef
{
stream = new SharpCompressStream(stream, bufferSize: Constants.BufferSize);
}
if (ZipArchive.IsZipFile(stream, password))
{
return true;

View File

@@ -40,14 +40,8 @@ public class DataDescriptorStream : Stream, IStreamStack
_start = _stream.Position;
_searchPosition = 0;
_done = false;
#if DEBUG_STREAMS
this.DebugConstruct(typeof(DataDescriptorStream));
#endif
}
internal bool IsRecording { get; private set; }
protected override void Dispose(bool disposing)
{
if (_isDisposed)
@@ -55,9 +49,6 @@ public class DataDescriptorStream : Stream, IStreamStack
return;
}
_isDisposed = true;
#if DEBUG_STREAMS
this.DebugDispose(typeof(DataDescriptorStream));
#endif
base.Dispose(disposing);
if (disposing)
{

View File

@@ -7,6 +7,10 @@ namespace SharpCompress.IO
{
internal partial class RewindableStream(Stream stream) : Stream
{
public static RewindableStream EnsureSeekable(Stream stream)
{
return stream.CanSeek ? new SeekableRewindableStream(stream) : new RewindableStream(stream);
}
private readonly int _bufferSize = Constants.RewindableBufferSize;
private byte[]? _buffer = ArrayPool<byte>.Shared.Rent(Constants.RewindableBufferSize);
private int _bufferLength = 0;
@@ -15,7 +19,7 @@ namespace SharpCompress.IO
private bool _isDisposed;
private long _streamPosition;
internal bool IsRecording { get; private set; }
internal virtual bool IsRecording { get; private set; }
protected override void Dispose(bool disposing)
{
@@ -33,44 +37,14 @@ namespace SharpCompress.IO
}
}
public void Rewind(bool stopRecording = false)
public virtual void Rewind(bool stopRecording = false)
{
_isBuffering = true;
IsRecording = !stopRecording;
_bufferPosition = 0;
}
public void Rewind(MemoryStream buffer)
{
long bufferLength = buffer.Length;
if (_bufferPosition >= bufferLength)
{
_bufferPosition -= (int)bufferLength;
}
else
{
int bytesToKeep = _bufferLength - _bufferPosition;
if (bytesToKeep > 0)
{
Array.Copy(_buffer!, _bufferPosition, _buffer!, 0, bytesToKeep);
}
if (bufferLength > _bufferSize)
{
throw new InvalidOperationException(
$"External buffer size ({bufferLength} bytes) exceeds internal buffer capacity ({_bufferSize} bytes)"
);
}
_bufferLength = (int)bufferLength;
_bufferPosition = 0;
buffer.Position = 0;
int bytesRead = buffer.Read(_buffer!, 0, _bufferLength);
_bufferLength = bytesRead;
_bufferPosition = 0;
}
_isBuffering = true;
}
public void StartRecording()
public virtual void StartRecording()
{
if (_bufferPosition != 0)
{
@@ -85,7 +59,7 @@ namespace SharpCompress.IO
IsRecording = true;
}
public void StopRecording()
public virtual void StopRecording()
{
_isBuffering = true;
IsRecording = false;

View File

@@ -0,0 +1,50 @@
using System;
using System.IO;
using System.Threading;
using System.Threading.Tasks;
namespace SharpCompress.IO;
internal sealed partial class SeekableRewindableStream
{
public override Task<int> ReadAsync(
byte[] buffer,
int offset,
int count,
CancellationToken cancellationToken
) => _underlyingStream.ReadAsync(buffer, offset, count, cancellationToken);
#if !LEGACY_DOTNET
public override ValueTask<int> ReadAsync(
Memory<byte> buffer,
CancellationToken cancellationToken = default
) => _underlyingStream.ReadAsync(buffer, cancellationToken);
public override ValueTask WriteAsync(
ReadOnlyMemory<byte> buffer,
CancellationToken cancellationToken = default
) => _underlyingStream.WriteAsync(buffer, cancellationToken);
public override ValueTask DisposeAsync()
{
_underlyingStream.Dispose();
return base.DisposeAsync();
}
#endif
public override Task WriteAsync(
byte[] buffer,
int offset,
int count,
CancellationToken cancellationToken
) => _underlyingStream.WriteAsync(buffer, offset, count, cancellationToken);
public override Task FlushAsync(CancellationToken cancellationToken) =>
_underlyingStream.FlushAsync(cancellationToken);
public override Task CopyToAsync(
Stream destination,
int bufferSize,
CancellationToken cancellationToken
) => _underlyingStream.CopyToAsync(destination, bufferSize, cancellationToken);
}

View File

@@ -0,0 +1,131 @@
using System;
using System.IO;
namespace SharpCompress.IO;
internal sealed partial class SeekableRewindableStream : RewindableStream
{
private readonly Stream _underlyingStream;
private long? _recordedPosition;
public SeekableRewindableStream(Stream stream)
: base(new NullStream())
{
if (stream is null)
{
throw new ArgumentNullException(nameof(stream));
}
if (!stream.CanSeek)
{
throw new ArgumentException("Stream must be seekable", nameof(stream));
}
_underlyingStream = stream;
}
public override bool CanRead => _underlyingStream.CanRead;
public override bool CanSeek => _underlyingStream.CanSeek;
public override bool CanWrite => _underlyingStream.CanWrite;
public override long Length => _underlyingStream.Length;
public override long Position
{
get => _underlyingStream.Position;
set => _underlyingStream.Position = value;
}
internal override bool IsRecording => _recordedPosition.HasValue;
public override void Flush() => _underlyingStream.Flush();
public override int Read(byte[] buffer, int offset, int count) =>
_underlyingStream.Read(buffer, offset, count);
#if !LEGACY_DOTNET
public override int Read(Span<byte> buffer) => _underlyingStream.Read(buffer);
#endif
public override long Seek(long offset, SeekOrigin origin) =>
_underlyingStream.Seek(offset, origin);
public override void SetLength(long value) => _underlyingStream.SetLength(value);
public override void Write(byte[] buffer, int offset, int count) =>
_underlyingStream.Write(buffer, offset, count);
#if !LEGACY_DOTNET
public override void Write(ReadOnlySpan<byte> buffer) => _underlyingStream.Write(buffer);
#endif
public override void Rewind(bool stopRecording = false)
{
if (!_recordedPosition.HasValue)
{
return;
}
_underlyingStream.Seek(_recordedPosition.Value, SeekOrigin.Begin);
if (stopRecording)
{
_recordedPosition = null;
}
}
public override void StartRecording()
{
_recordedPosition = _underlyingStream.Position;
}
public override void StopRecording()
{
_recordedPosition = null;
}
protected override void Dispose(bool disposing)
{
if (disposing)
{
_underlyingStream.Dispose();
}
base.Dispose(disposing);
}
private sealed class NullStream : Stream
{
public override bool CanRead => true;
public override bool CanSeek => false;
public override bool CanWrite => false;
public override long Length => throw new NotSupportedException();
public override long Position
{
get => throw new NotSupportedException();
set => throw new NotSupportedException();
}
public override void Flush() { }
public override int Read(byte[] buffer, int offset, int count) => 0;
#if !LEGACY_DOTNET
public override int Read(Span<byte> buffer) => 0;
#endif
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();
#if !LEGACY_DOTNET
public override void Write(ReadOnlySpan<byte> buffer) => throw new NotSupportedException();
#endif
}
}

View File

@@ -34,7 +34,7 @@ public static partial class ReaderFactory
stream.NotNull(nameof(stream));
options ??= new ReaderOptions() { LeaveStreamOpen = false };
var bStream = new RewindableStream(stream);
var bStream = RewindableStream.EnsureSeekable(stream);
bStream.StartRecording();
var factories = Factories.Factory.Factories.OfType<Factories.Factory>();
@@ -52,7 +52,7 @@ public static partial class ReaderFactory
&& reader != null
)
{
bStream.StopRecording();
bStream.Rewind(true);
return reader;
}
}
@@ -66,6 +66,7 @@ public static partial class ReaderFactory
bStream.Rewind();
if (factory.TryOpenReader(bStream, options, out var reader) && reader != null)
{
bStream.Rewind(true);
return reader;
}
}
@@ -74,6 +75,4 @@ public static partial class ReaderFactory
"Cannot determine compressed stream type. Supported Reader Formats: Ace, Arc, Arj, Zip, GZip, BZip2, Tar, Rar, LZip, XZ, ZStandard"
);
}
// Async methods moved to ReaderFactory.Async.cs
}

View File

@@ -465,4 +465,152 @@ public class RewindableStreamAsyncTest
Assert.False(stream.IsRecording);
}
[Fact]
public async ValueTask TestStopRecordingThenRewindAsync()
{
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.Write(8);
bw.Flush();
ms.Position = 0;
var stream = new RewindableStream(ms);
stream.StartRecording();
// Read first 4 values (gets buffered)
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));
// Stop recording
stream.StopRecording();
Assert.False(stream.IsRecording);
// Rewind to start of buffer
stream.Rewind(true);
// Should be able to read from buffer again
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));
// Continue reading remaining data from underlying stream
Assert.Equal(5, await ReadInt32Async(stream).ConfigureAwait(false));
Assert.Equal(6, await ReadInt32Async(stream).ConfigureAwait(false));
Assert.Equal(7, await ReadInt32Async(stream).ConfigureAwait(false));
Assert.Equal(8, await ReadInt32Async(stream).ConfigureAwait(false));
}
[Fact]
public async ValueTask TestMultipleRewindsAfterStopRecordingAsync()
{
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.Write(8);
bw.Flush();
ms.Position = 0;
var stream = new RewindableStream(ms);
stream.StartRecording();
// Read first 4 values (gets buffered)
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));
// Stop recording
stream.StopRecording();
Assert.False(stream.IsRecording);
// First rewind - read all buffered data, then continue with underlying stream
stream.Rewind();
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));
Assert.Equal(5, await ReadInt32Async(stream).ConfigureAwait(false));
Assert.Equal(6, await ReadInt32Async(stream).ConfigureAwait(false));
// Second rewind - should still be able to read from buffer
stream.Rewind();
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));
// Third rewind - still works
stream.Rewind();
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));
// Continue reading from underlying stream (values 7, 8 since 5, 6 were already consumed)
Assert.Equal(7, await ReadInt32Async(stream).ConfigureAwait(false));
Assert.Equal(8, await ReadInt32Async(stream).ConfigureAwait(false));
}
[Fact]
public async ValueTask TestStopRecordingTwiceThrowsAsync()
{
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.StartRecording();
Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false));
// First StopRecording should succeed
stream.StopRecording();
Assert.False(stream.IsRecording);
// Second StopRecording should throw
Assert.Throws<InvalidOperationException>(() => stream.StopRecording());
}
[Fact]
public async ValueTask TestStartRecordingAfterStopRecordingThrowsAsync()
{
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.StartRecording();
Assert.Equal(1, await ReadInt32Async(stream).ConfigureAwait(false));
// Stop recording
stream.StopRecording();
Assert.False(stream.IsRecording);
// Trying to start recording again should throw
Assert.Throws<InvalidOperationException>(() => stream.StartRecording());
}
}

View File

@@ -94,47 +94,6 @@ public class RewindableStreamTest
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()
{
@@ -482,6 +441,212 @@ public class RewindableStreamTest
Assert.False(stream.IsRecording);
}
[Fact]
public void TestStopRecordingThenRewind()
{
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.Write(8);
bw.Flush();
ms.Position = 0;
var stream = new RewindableStream(ms);
stream.StartRecording();
var br = new BinaryReader(new ForwardOnlyStream(stream));
// Read first 4 values (gets buffered)
Assert.Equal(1, br.ReadInt32());
Assert.Equal(2, br.ReadInt32());
Assert.Equal(3, br.ReadInt32());
Assert.Equal(4, br.ReadInt32());
// Stop recording
stream.StopRecording();
Assert.False(stream.IsRecording);
// Rewind to start of buffer
stream.Rewind(true);
// Should be able to read from buffer again
Assert.Equal(1, br.ReadInt32());
Assert.Equal(2, br.ReadInt32());
Assert.Equal(3, br.ReadInt32());
Assert.Equal(4, br.ReadInt32());
// Rewind to start of buffer
stream.Rewind();
// Should be able to read from buffer again
Assert.Equal(1, br.ReadInt32());
Assert.Equal(2, br.ReadInt32());
Assert.Equal(3, br.ReadInt32());
Assert.Equal(4, br.ReadInt32());
// Continue reading remaining data from underlying stream
Assert.Equal(5, br.ReadInt32());
Assert.Equal(6, br.ReadInt32());
Assert.Equal(7, br.ReadInt32());
Assert.Equal(8, br.ReadInt32());
}
[Fact]
public void TestNonSeekableStream_StopRecordingThenRewind()
{
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.Write(8);
bw.Flush();
ms.Position = 0;
var nonSeekableStream = new NonSeekableStreamWrapper(ms);
var stream = new RewindableStream(nonSeekableStream);
stream.StartRecording();
var br = new BinaryReader(stream);
// Read first 4 values (gets buffered)
Assert.Equal(1, br.ReadInt32());
Assert.Equal(2, br.ReadInt32());
Assert.Equal(3, br.ReadInt32());
Assert.Equal(4, br.ReadInt32());
// Stop recording
stream.StopRecording();
Assert.False(stream.IsRecording);
// Rewind to start of buffer
stream.Rewind(true);
// Should be able to read from buffer again
Assert.Equal(1, br.ReadInt32());
Assert.Equal(2, br.ReadInt32());
Assert.Equal(3, br.ReadInt32());
Assert.Equal(4, br.ReadInt32());
// Continue reading remaining data from underlying stream
Assert.Equal(5, br.ReadInt32());
Assert.Equal(6, br.ReadInt32());
Assert.Equal(7, br.ReadInt32());
Assert.Equal(8, br.ReadInt32());
}
[Fact]
public void TestMultipleRewindsAfterStopRecording()
{
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.Write(8);
bw.Flush();
ms.Position = 0;
var stream = new RewindableStream(ms);
stream.StartRecording();
var br = new BinaryReader(stream);
// Read first 4 values (gets buffered)
Assert.Equal(1, br.ReadInt32());
Assert.Equal(2, br.ReadInt32());
Assert.Equal(3, br.ReadInt32());
Assert.Equal(4, br.ReadInt32());
// Stop recording
stream.StopRecording();
Assert.False(stream.IsRecording);
// First rewind - read all buffered data, then continue with underlying stream
stream.Rewind();
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());
// Second rewind - should still be able to read from buffer
stream.Rewind();
Assert.Equal(1, br.ReadInt32());
Assert.Equal(2, br.ReadInt32());
Assert.Equal(3, br.ReadInt32());
Assert.Equal(4, br.ReadInt32());
// Third rewind - still works
stream.Rewind();
Assert.Equal(1, br.ReadInt32());
Assert.Equal(2, br.ReadInt32());
Assert.Equal(3, br.ReadInt32());
Assert.Equal(4, br.ReadInt32());
// Continue reading from underlying stream (values 7, 8 since 5, 6 were already consumed)
Assert.Equal(7, br.ReadInt32());
Assert.Equal(8, br.ReadInt32());
}
[Fact]
public void TestStopRecordingTwiceThrows()
{
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.StartRecording();
var br = new BinaryReader(stream);
Assert.Equal(1, br.ReadInt32());
// First StopRecording should succeed
stream.StopRecording();
Assert.False(stream.IsRecording);
// Second StopRecording should throw
Assert.Throws<InvalidOperationException>(() => stream.StopRecording());
}
[Fact]
public void TestStartRecordingAfterStopRecordingThrows()
{
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.StartRecording();
var br = new BinaryReader(stream);
Assert.Equal(1, br.ReadInt32());
// Stop recording
stream.StopRecording();
Assert.False(stream.IsRecording);
// Trying to start recording again should throw
Assert.Throws<InvalidOperationException>(() => stream.StartRecording());
}
private class NonSeekableStreamWrapper : Stream
{
private readonly Stream _baseStream;

View File

@@ -0,0 +1,219 @@
using System;
using System.IO;
using System.Linq;
using System.Threading;
using System.Threading.Tasks;
using SharpCompress.IO;
using Xunit;
namespace SharpCompress.Test.Streams;
public class SeekableRewindableStreamAsyncTest
{
[Fact]
public async Task ReadAsync_Buffers()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[5];
int bytesRead = await stream.ReadAsync(buffer, 0, buffer.Length).ConfigureAwait(false);
Assert.Equal(5, bytesRead);
Assert.Equal(new byte[] { 1, 2, 3, 4, 5 }, buffer);
}
[Fact]
public async Task ReadAsync_WithCancellation()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[5];
var cts = new CancellationTokenSource();
int bytesRead = await stream
.ReadAsync(buffer, 0, buffer.Length, cts.Token)
.ConfigureAwait(false);
Assert.Equal(5, bytesRead);
Assert.Equal(new byte[] { 1, 2, 3, 4, 5 }, buffer);
}
[Fact]
public async Task ReadAsync_PartialRead()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[10];
int bytesRead = await stream.ReadAsync(buffer, 0, 10).ConfigureAwait(false);
Assert.Equal(5, bytesRead);
Assert.Equal(new byte[] { 1, 2, 3, 4, 5 }, buffer.Take(5).ToArray());
}
[Fact]
public async Task WriteAsync_Buffers()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var data = new byte[] { 1, 2, 3, 4, 5 };
await stream.WriteAsync(data, 0, data.Length).ConfigureAwait(false);
Assert.Equal(data, ms.ToArray());
}
[Fact]
public async Task WriteAsync_WithCancellation()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var data = new byte[] { 1, 2, 3, 4, 5 };
var cts = new CancellationTokenSource();
await stream.WriteAsync(data, 0, data.Length, cts.Token).ConfigureAwait(false);
Assert.Equal(data, ms.ToArray());
}
[Fact]
public async Task FlushAsync_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var data = new byte[] { 1, 2, 3, 4, 5 };
await stream.WriteAsync(data, 0, data.Length).ConfigureAwait(false);
await stream.FlushAsync().ConfigureAwait(false);
Assert.Equal(5, ms.Length);
}
[Fact]
public async Task CopyToAsync_CopiesAllData()
{
var sourceMs = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(sourceMs);
var destinationMs = new MemoryStream();
await stream.CopyToAsync(destinationMs, 4096).ConfigureAwait(false);
Assert.Equal(new byte[] { 1, 2, 3, 4, 5 }, destinationMs.ToArray());
}
[Fact]
public async Task ReadAsyncAndSeek_MultipleOperations()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5, 6, 7, 8, 9, 10 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[3];
await stream.ReadAsync(buffer, 0, 3).ConfigureAwait(false);
Assert.Equal(new byte[] { 1, 2, 3 }, buffer);
Assert.Equal(3, stream.Position);
stream.Seek(7, SeekOrigin.Begin);
Assert.Equal(7, stream.Position);
Array.Clear(buffer, 0, buffer.Length);
await stream.ReadAsync(buffer, 0, 2).ConfigureAwait(false);
Assert.Equal(new byte[] { 8, 9, 0 }, buffer);
Assert.Equal(9, stream.Position);
}
[Fact]
public async Task WriteAsyncAndReadAsync_WrittenDataIsReadable()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var writeData = new byte[] { 1, 2, 3, 4, 5 };
await stream.WriteAsync(writeData, 0, writeData.Length).ConfigureAwait(false);
stream.Position = 0;
var readBuffer = new byte[5];
await stream.ReadAsync(readBuffer, 0, 5).ConfigureAwait(false);
Assert.Equal(writeData, readBuffer);
}
[Fact]
public async Task AsyncOperationsDoNotCauseRecording()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
stream.StartRecording();
var buffer = new byte[3];
await stream.ReadAsync(buffer, 0, 3).ConfigureAwait(false);
Assert.Equal(new byte[] { 1, 2, 3 }, buffer);
Assert.Equal(3, stream.Position);
stream.Rewind(true);
Assert.Equal(3, stream.Position);
var buffer2 = new byte[2];
await stream.ReadAsync(buffer2, 0, 2).ConfigureAwait(false);
Assert.Equal(new byte[] { 4, 5 }, buffer2);
Assert.Equal(5, stream.Position);
}
}
#if !LEGACY_DOTNET
public partial class SeekableRewindableStreamMemoryAsyncTest
{
[Fact]
public async ValueTask ReadAsync_Memory()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[5];
int bytesRead = await stream.ReadAsync(buffer).ConfigureAwait(false);
Assert.Equal(5, bytesRead);
Assert.Equal(new byte[] { 1, 2, 3, 4, 5 }, buffer);
}
[Fact]
public async ValueTask ReadAsync_Memory_WithCancellation()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[5];
var cts = new CancellationTokenSource();
int bytesRead = await stream.ReadAsync(buffer, cts.Token).ConfigureAwait(false);
Assert.Equal(5, bytesRead);
Assert.Equal(new byte[] { 1, 2, 3, 4, 5 }, buffer);
}
[Fact]
public async ValueTask WriteAsync_Memory()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var data = new byte[] { 1, 2, 3, 4, 5 };
await stream.WriteAsync(data).ConfigureAwait(false);
Assert.Equal(data, ms.ToArray());
}
[Fact]
public async ValueTask WriteAsync_Memory_WithCancellation()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var data = new byte[] { 1, 2, 3, 4, 5 };
var cts = new CancellationTokenSource();
await stream.WriteAsync(data, cts.Token).ConfigureAwait(false);
Assert.Equal(data, ms.ToArray());
}
[Fact]
public async ValueTask ReadMemoryAndWriteMemory_MemoryOperations()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var writeData = new byte[] { 1, 2, 3, 4, 5 };
await stream.WriteAsync(writeData).ConfigureAwait(false);
stream.Position = 0;
var readBuffer = new byte[5];
await stream.ReadAsync(readBuffer).ConfigureAwait(false);
Assert.Equal(writeData, readBuffer);
}
[Fact]
public async ValueTask DisposeAsync_DisposesUnderlyingStream()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
await stream.DisposeAsync().ConfigureAwait(false);
Assert.Throws<ObjectDisposedException>(() => ms.Read(new byte[1], 0, 1));
}
}
#endif

View File

@@ -0,0 +1,284 @@
using System;
using System.IO;
using System.Linq;
using SharpCompress.IO;
using SharpCompress.Test.Mocks;
using Xunit;
namespace SharpCompress.Test.Streams;
public class SeekableRewindableStreamTest
{
[Fact]
public void Constructor_ThrowsOnNullStream()
{
Assert.Throws<ArgumentNullException>(() => new SeekableRewindableStream(null!));
}
[Fact]
public void Constructor_ThrowsOnNonSeekableStream()
{
var nonSeekable = new ForwardOnlyStream(new MemoryStream());
Assert.Throws<ArgumentException>(() => new SeekableRewindableStream(nonSeekable));
}
[Fact]
public void Constructor_AcceptsSeekableStream()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
Assert.NotNull(stream);
}
[Fact]
public void CanRead_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
Assert.Equal(ms.CanRead, stream.CanRead);
}
[Fact]
public void CanSeek_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
Assert.Equal(ms.CanSeek, stream.CanSeek);
}
[Fact]
public void CanWrite_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
Assert.Equal(ms.CanWrite, stream.CanWrite);
}
[Fact]
public void Length_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
Assert.Equal(5, stream.Length);
}
[Fact]
public void Position_Getter_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
ms.Position = 2;
var stream = new SeekableRewindableStream(ms);
Assert.Equal(2, stream.Position);
}
[Fact]
public void Position_Setter_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
stream.Position = 3;
Assert.Equal(3, ms.Position);
Assert.Equal(3, stream.Position);
}
[Fact]
public void IsRecording_AlwaysFalse()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
Assert.False(stream.IsRecording);
}
[Fact]
public void Read_Buffers()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[5];
int bytesRead = stream.Read(buffer, 0, buffer.Length);
Assert.Equal(5, bytesRead);
Assert.Equal(new byte[] { 1, 2, 3, 4, 5 }, buffer);
}
[Fact]
public void Seek_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
long result = stream.Seek(3, SeekOrigin.Begin);
Assert.Equal(3, result);
Assert.Equal(3, ms.Position);
}
[Fact]
public void SetLength_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
stream.SetLength(20);
Assert.Equal(20, stream.Length);
Assert.Equal(20, ms.Length);
}
[Fact]
public void Write_DelegatesToUnderlyingStream()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var data = new byte[] { 1, 2, 3, 4, 5 };
stream.Write(data, 0, data.Length);
Assert.Equal(data, ms.ToArray());
}
[Fact]
public void Rewind_IsNoOp()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
stream.Rewind();
Assert.Equal(0, stream.Position);
ms.Position = 2;
stream.Rewind(true);
Assert.Equal(2, stream.Position);
}
[Fact]
public void StartRecording_IsNoOp()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
stream.StartRecording();
Assert.False(stream.IsRecording);
}
[Fact]
public void StopRecording_IsNoOp()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
stream.StopRecording();
Assert.False(stream.IsRecording);
}
[Fact]
public void Dispose_DisposesUnderlyingStream()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
stream.Dispose();
Assert.Throws<ObjectDisposedException>(() => ms.Read(new byte[1], 0, 1));
}
[Fact]
public void ReadAndSeek_MultipleOperations()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5, 6, 7, 8, 9, 10 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[3];
stream.Read(buffer, 0, 3);
Assert.Equal(new byte[] { 1, 2, 3 }, buffer);
Assert.Equal(3, stream.Position);
stream.Seek(7, SeekOrigin.Begin);
Assert.Equal(7, stream.Position);
Array.Clear(buffer, 0, buffer.Length);
stream.Read(buffer, 0, 2);
Assert.Equal(new byte[] { 8, 9, 0 }, buffer);
Assert.Equal(9, stream.Position);
}
[Fact]
public void SeekWithDifferentOrigins()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5, 6, 7, 8, 9, 10 });
var stream = new SeekableRewindableStream(ms);
stream.Seek(3, SeekOrigin.Begin);
Assert.Equal(3, stream.Position);
stream.Seek(2, SeekOrigin.Current);
Assert.Equal(5, stream.Position);
stream.Seek(-3, SeekOrigin.End);
Assert.Equal(7, stream.Position);
}
[Fact]
public void WriteAndRead_WrittenDataIsReadable()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var writeData = new byte[] { 1, 2, 3, 4, 5 };
stream.Write(writeData, 0, writeData.Length);
stream.Position = 0;
var readBuffer = new byte[5];
stream.Read(readBuffer, 0, 5);
Assert.Equal(writeData, readBuffer);
}
[Fact]
public void RecordingOperationsDoNotAffectStream()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
stream.StartRecording();
var buffer = new byte[3];
stream.Read(buffer, 0, 3);
Assert.Equal(new byte[] { 1, 2, 3 }, buffer);
Assert.Equal(3, stream.Position);
stream.Rewind(true);
Assert.Equal(3, stream.Position);
var buffer2 = new byte[2];
stream.Read(buffer2, 0, 2);
Assert.Equal(new byte[] { 4, 5 }, buffer2);
Assert.Equal(5, stream.Position);
}
}
#if !LEGACY_DOTNET
public partial class SeekableRewindableSpanTest
{
[Fact]
public void Read_Span()
{
var ms = new MemoryStream(new byte[] { 1, 2, 3, 4, 5 });
var stream = new SeekableRewindableStream(ms);
var buffer = new byte[5];
int bytesRead = stream.Read(buffer);
Assert.Equal(5, bytesRead);
Assert.Equal(new byte[] { 1, 2, 3, 4, 5 }, buffer);
}
[Fact]
public void Write_Span()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var data = new byte[] { 1, 2, 3, 4, 5 };
stream.Write(data);
Assert.Equal(data, ms.ToArray());
}
[Fact]
public void ReadAndWrite_SpanOperations()
{
var ms = new MemoryStream();
var stream = new SeekableRewindableStream(ms);
var writeData = new byte[] { 1, 2, 3, 4, 5 };
stream.Write(writeData);
stream.Position = 0;
var readBuffer = new byte[5];
stream.Read(readBuffer);
Assert.Equal(writeData, readBuffer);
}
}
#endif