Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 104 additions & 0 deletions internal/nzbfilesystem/aes_tail_read_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
package nzbfilesystem

import (
"bytes"
"context"
cryptoaes "crypto/aes"
"crypto/cipher"
"errors"
"io"
"testing"
"time"

"github.com/kipsilabs/altmount/internal/encryption/aes"
metapb "github.com/kipsilabs/altmount/internal/metadata/proto"
"github.com/kipsilabs/altmount/internal/testsupport/fakepool"
"github.com/kipsilabs/altmount/internal/testsupport/segments"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

// newAESTestMVF builds an AES-CBC encrypted file whose plaintext length is
// `pad` bytes short of a block boundary, so the final ciphertext block extends
// past FileSize. The segments cover the whole padded ciphertext (as the RAR5
// and 7z importers store it) unless the caller trims them afterwards.
func newAESTestMVF(t *testing.T, n, segSize, pad int) (*MetadataVirtualFile, []byte) {
t.Helper()
total := n * segSize
plainLen := total - pad
plain := segments.FileBytes(n, segSize)[:plainLen]

key := bytes.Repeat([]byte{0x42}, 32)
iv := bytes.Repeat([]byte{0x24}, 16)
block, err := cryptoaes.NewCipher(key)
require.NoError(t, err)
padded := append(append([]byte{}, plain...), bytes.Repeat([]byte{byte(pad)}, pad)...)
ct := make([]byte, total)
cipher.NewCBCEncrypter(block, iv).CryptBlocks(ct, padded)

fp := fakepool.New()
for i := range n {
fp.SetBehavior(segments.MessageID(i), fakepool.SegmentBehavior{Bytes: ct[i*segSize : (i+1)*segSize]})
}
mvf := newTestMVF(t, context.Background(), fp, n, segSize, 4)
mvf.meta.FileSize = int64(plainLen)
mvf.meta.Encryption = metapb.Encryption_AES
mvf.meta.AesKey = key
mvf.meta.AesIv = iv
mvf.aesCipher = aes.NewAesCipher()
return mvf, plain
}

// readAllOrHang fails the test instead of hanging forever when the read path
// spins: the bug this guards against held mvf.mu in a busy loop.
func readAllOrHang(t *testing.T, r io.Reader) ([]byte, error) {
t.Helper()
type res struct {
b []byte
err error
}
done := make(chan res, 1)
go func() {
b, err := io.ReadAll(r)
done <- res{b, err}
}()
select {
case r := <-done:
return r.b, r.err
case <-time.After(10 * time.Second):
t.Fatal("read hung: the reader is spinning instead of returning")
return nil, nil
}
}

func TestAESTailReadReachesPaddedFinalBlock(t *testing.T) {
const n, segSize, pad = 8, 64 << 10, 8
mvf, plain := newAESTestMVF(t, n, segSize, pad)
start := int64(len(plain) - 2000)

_, err := mvf.Seek(start, io.SeekStart)
require.NoError(t, err)

got, err := readAllOrHang(t, mvf)
require.NoError(t, err)
assert.True(t, bytes.Equal(got, plain[start:]), "got %d bytes, want %d", len(got), len(plain)-int(start))
}

func TestReadReturnsErrorWhenReaderCannotReachRangeEnd(t *testing.T) {
const n, segSize, pad = 8, 64 << 10, 8
mvf, plain := newAESTestMVF(t, n, segSize, pad)
// Truncated metadata: the padded final block is not addressable, so the
// decryptor can never produce the last bytes. That must surface as an
// error, never as a retry loop.
last := mvf.meta.SegmentData[n-1]
last.EndOffset -= pad
start := int64(len(plain) - 2000)

_, err := mvf.Seek(start, io.SeekStart)
require.NoError(t, err)

got, err := readAllOrHang(t, mvf)
require.Error(t, err)
assert.True(t, errors.Is(err, io.ErrUnexpectedEOF), "want ErrUnexpectedEOF, got %v", err)
assert.True(t, bytes.HasPrefix(plain[start:], got), "delivered bytes must be a prefix of the plaintext")
}
44 changes: 38 additions & 6 deletions internal/nzbfilesystem/metadata_remote_file.go
Original file line number Diff line number Diff line change
Expand Up @@ -1139,6 +1139,15 @@ func (idx *segmentOffsetIndex) findSegmentForOffset(offset int64) int {
return lo - 1
}

// totalBytes is the number of bytes the indexed segments cover.
func (idx *segmentOffsetIndex) totalBytes() int64 {
if idx == nil || len(idx.offsets) == 0 {
return 0
}
n := len(idx.offsets)
return idx.offsets[n-1] + idx.sizes[n-1]
}

// getOffsetForSegment returns the cumulative file offset at the start of the given segment index
// Returns 0 if the index is invalid or out of bounds
func (idx *segmentOffsetIndex) getOffsetForSegment(segmentIndex int) int64 {
Expand Down Expand Up @@ -1177,6 +1186,7 @@ func (mvf *MetadataVirtualFile) Read(p []byte) (n int, err error) {
return 0, ErrFileClosed
}

stalledAt := int64(-1)
for n < len(p) {
if err := mvf.ensureReader(); err != nil {
return n, err
Expand All @@ -1198,6 +1208,13 @@ func (mvf *MetadataVirtualFile) Read(p []byte) (n int, err error) {

if readErr != nil {
if errors.Is(readErr, io.EOF) && mvf.hasMoreDataToRead() {
// A rebuilt reader that ends at the same offset again cannot reach
// the range end; rotating once more would spin here holding mvf.mu.
if totalRead == 0 && mvf.position == stalledAt {
mvf.closeCurrentReader()
return n, fmt.Errorf("%w: reader ended at offset %d before the requested end", io.ErrUnexpectedEOF, mvf.position)
}
stalledAt = mvf.position
// Close current reader and try to get a new one for the next range in next iteration
mvf.closeCurrentReader()
continue
Expand Down Expand Up @@ -1327,6 +1344,7 @@ func (mvf *MetadataVirtualFile) ReadAtContext(readCtx context.Context, p []byte,
}
buf := p[:want]
var sharedErr error
stalledAt := int64(-1)
for n < int(want) {
rn, readErr := mvf.reader.Read(buf[n:])
n += rn
Expand All @@ -1341,6 +1359,14 @@ func (mvf *MetadataVirtualFile) ReadAtContext(readCtx context.Context, p []byte,

if readErr != nil {
if errors.Is(readErr, io.EOF) && mvf.hasMoreDataToRead() {
// Same no-progress guard as Read: a rebuilt reader ending at the
// same offset cannot reach the range end.
at := off + int64(n)
if rn == 0 && at == stalledAt {
sharedErr = fmt.Errorf("%w: reader ended at offset %d before the requested end", io.ErrUnexpectedEOF, at)
break
}
stalledAt = at
mvf.closeCurrentReader()
if rotateErr := mvf.ensureReader(); rotateErr != nil {
sharedErr = rotateErr
Expand Down Expand Up @@ -2025,18 +2051,24 @@ func (mvf *MetadataVirtualFile) createUsenetReader(ctx context.Context, start, e
if len(mvf.meta.SegmentData) == 0 {
return nil, ErrMissmatchedSegments
}
if start >= mvf.meta.FileSize {
return nil, io.EOF
}
if end >= mvf.meta.FileSize {
end = mvf.meta.FileSize - 1
}

// Build segment offset index lazily on first read (thread-safe via sync.Once)
mvf.segmentIndexOnce.Do(func() {
mvf.segmentIndex = buildSegmentIndex(mvf.meta.SegmentData)
})

// Bound the range by what the segments cover, not by FileSize: an
// AES-encrypted file's segments extend up to 15 bytes past FileSize (the
// padded final block), and the decryptor needs them to produce the last
// plaintext bytes.
covered := mvf.segmentIndex.totalBytes()
if start >= covered {
return nil, io.EOF
}
if end >= covered {
end = covered - 1
}

loader := newMetadataSegmentLoader(mvf.meta.SegmentData)

// segmentIndex is always non-nil here (built by segmentIndexOnce.Do above).
Expand Down
Loading