diff options
| author | Jeff Halter <868228+jhalter@users.noreply.github.com> | 2026-07-10 09:48:18 -0700 |
|---|---|---|
| committer | Jeff Halter <868228+jhalter@users.noreply.github.com> | 2026-07-10 09:48:18 -0700 |
| commit | 21f24d24fd6f501b32f15a2bef41c89cc461f623 (patch) | |
| tree | ba4952fb82f3ef9c265228633f27536866e23931 /hotline | |
| parent | ae44fb222ec73cae8441f5a5d9a21f8585fe587f (diff) | |
Overhaul regression testing: e2e suite, fuzzing, CI, and bug fixes
Add a protocol-level end-to-end suite (in-process fully wired server on
an ephemeral port pair, driven by hotline.Client over TCP) covering
handshake, login, public and private chat, message board, threaded
news, file list/download/upload, account admin, disconnect
notification, and shutdown broadcast. The harness retries on a fresh
port pair when another process steals a probed port before
ListenAndServe binds it, and Server gains WithConnectionRateLimit so
tests can disable the per-IP connection throttle.
Add native fuzz tests for Transaction, Field, and flattened file
object decoding, and fix the bugs the new tests surfaced:
- Transaction.Write panicked on out-of-range attacker-controlled size
fields, and transactionScanner's uint32 length addition could wrap
and yield a truncated token. The information fork size declared in
an untrusted fork header is now bounded too.
- Client keepalive read c.done unsynchronized while Disconnect
replaces it under the mutex.
- The shared Agreement's Seek+ReadAll login path raced concurrent
logins; the server now prefers an AgreementBytes() snapshot.
Fill unit-test gaps (main's config-copy helpers, file resume data,
ReloaderFunc, R2 error injection and env validation) and add a CI test
workflow (build/vet + race-enabled shuffled suite), fixed lint
workflow triggers with golangci-lint v2.6, and Makefile test/cover/
lint/fuzz targets.
Diffstat (limited to 'hotline')
| -rw-r--r-- | hotline/client.go | 12 | ||||
| -rw-r--r-- | hotline/client_test.go | 24 | ||||
| -rw-r--r-- | hotline/field_fuzz_test.go | 76 | ||||
| -rw-r--r-- | hotline/file_resume_data_test.go | 87 | ||||
| -rw-r--r-- | hotline/flattened_file_object.go | 7 | ||||
| -rw-r--r-- | hotline/flattened_file_object_fuzz_test.go | 95 | ||||
| -rw-r--r-- | hotline/r2_file_store.go | 10 | ||||
| -rw-r--r-- | hotline/r2_file_store_test.go | 113 | ||||
| -rw-r--r-- | hotline/server.go | 32 | ||||
| -rw-r--r-- | hotline/server_blackbox_test.go | 30 | ||||
| -rw-r--r-- | hotline/transaction.go | 15 | ||||
| -rw-r--r-- | hotline/transaction_fuzz_test.go | 149 |
12 files changed, 610 insertions, 40 deletions
diff --git a/hotline/client.go b/hotline/client.go index 23bfe9c..b09a01d 100644 --- a/hotline/client.go +++ b/hotline/client.go @@ -84,15 +84,19 @@ func (c *Client) Connect(address, login, passwd string) (err error) { return fmt.Errorf("error sending login transaction: %w", err) } - // start keepalive go routine - go func() { _ = c.keepalive() }() + // start keepalive go routine. Capture the done channel now so keepalive never races with + // Disconnect, which replaces c.done under the mutex. + c.mu.Lock() + done := c.done + c.mu.Unlock() + go func() { _ = c.keepalive(done) }() return nil } const keepaliveInterval = 300 * time.Second -func (c *Client) keepalive() error { +func (c *Client) keepalive(done <-chan struct{}) error { ticker := time.NewTicker(keepaliveInterval) defer ticker.Stop() @@ -100,7 +104,7 @@ func (c *Client) keepalive() error { select { case <-ticker.C: _ = c.Send(NewTransaction(TranKeepAlive, [2]byte{})) - case <-c.done: + case <-done: return nil } } diff --git a/hotline/client_test.go b/hotline/client_test.go index 23b6155..fb5a3a9 100644 --- a/hotline/client_test.go +++ b/hotline/client_test.go @@ -21,8 +21,8 @@ func newTestClient() *Client { func TestClient_Handshake(t *testing.T) { t.Run("successful handshake", func(t *testing.T) { clientConn, serverConn := net.Pipe() - defer clientConn.Close() - defer serverConn.Close() + defer func() { _ = clientConn.Close() }() + defer func() { _ = serverConn.Close() }() c := newTestClient() c.Connection = clientConn @@ -41,8 +41,8 @@ func TestClient_Handshake(t *testing.T) { t.Run("server returns error response", func(t *testing.T) { clientConn, serverConn := net.Pipe() - defer clientConn.Close() - defer serverConn.Close() + defer func() { _ = clientConn.Close() }() + defer func() { _ = serverConn.Close() }() c := newTestClient() c.Connection = clientConn @@ -61,7 +61,7 @@ func TestClient_Handshake(t *testing.T) { t.Run("connection closed during read", func(t *testing.T) { clientConn, serverConn := net.Pipe() - defer clientConn.Close() + defer func() { _ = clientConn.Close() }() c := newTestClient() c.Connection = clientConn @@ -69,7 +69,7 @@ func TestClient_Handshake(t *testing.T) { go func() { buf := make([]byte, 12) _, _ = io.ReadFull(serverConn, buf) - serverConn.Close() + _ = serverConn.Close() }() err := c.Handshake() @@ -81,8 +81,8 @@ func TestClient_Handshake(t *testing.T) { func TestClient_Send(t *testing.T) { t.Run("sends non-reply transaction and tracks it", func(t *testing.T) { clientConn, serverConn := net.Pipe() - defer clientConn.Close() - defer serverConn.Close() + defer func() { _ = clientConn.Close() }() + defer func() { _ = serverConn.Close() }() c := newTestClient() c.Connection = clientConn @@ -104,8 +104,8 @@ func TestClient_Send(t *testing.T) { t.Run("reply transactions are not tracked", func(t *testing.T) { clientConn, serverConn := net.Pipe() - defer clientConn.Close() - defer serverConn.Close() + defer func() { _ = clientConn.Close() }() + defer func() { _ = serverConn.Close() }() c := newTestClient() c.Connection = clientConn @@ -195,7 +195,7 @@ func TestClient_HandleTransaction(t *testing.T) { func TestClient_Disconnect(t *testing.T) { t.Run("closes connection and done channel", func(t *testing.T) { clientConn, serverConn := net.Pipe() - defer serverConn.Close() + defer func() { _ = serverConn.Close() }() c := newTestClient() c.Connection = clientConn @@ -220,7 +220,7 @@ func TestClient_Disconnect(t *testing.T) { t.Run("handles nil done channel", func(t *testing.T) { clientConn, serverConn := net.Pipe() - defer serverConn.Close() + defer func() { _ = serverConn.Close() }() c := newTestClient() c.Connection = clientConn diff --git a/hotline/field_fuzz_test.go b/hotline/field_fuzz_test.go new file mode 100644 index 0000000..1af4f5e --- /dev/null +++ b/hotline/field_fuzz_test.go @@ -0,0 +1,76 @@ +package hotline + +import ( + "io" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// FuzzFieldScanner verifies the field-framing split func never panics or advances past its input. +func FuzzFieldScanner(f *testing.F) { + f.Add([]byte{0x00, 0x65, 0x00, 0x03, 0x68, 0x61, 0x69}) // FieldData "hai" + f.Add([]byte{0x00, 0x65, 0xff, 0xff}) // declared size larger than input + f.Add([]byte{0x00}) // shorter than the size field + + f.Fuzz(func(t *testing.T, data []byte) { + advance, token, err := FieldScanner(data, false) + if err != nil { + return + } + if advance == 0 { + assert.Nil(t, token, "no advance must produce no token") + return + } + require.LessOrEqual(t, advance, len(data), "scanner advanced past its input") + require.Len(t, token, advance, "token length must match advance") + }) +} + +// FuzzFieldWrite feeds raw untrusted bytes to the Field decoder. If decoding succeeds, +// re-encoding must reproduce exactly the bytes that were consumed. +func FuzzFieldWrite(f *testing.F) { + f.Add([]byte{0x00, 0x65, 0x00, 0x03, 0x68, 0x61, 0x69}) // FieldData "hai" + f.Add([]byte{0x00, 0x65, 0x00, 0x00}) // empty data + f.Add([]byte{0x00, 0x65, 0xff, 0xff, 0x00}) // declared size overruns buffer + + f.Fuzz(func(t *testing.T, data []byte) { + var field Field + n, err := field.Write(data) + if err != nil { + return + } + + encoded, err := io.ReadAll(&field) + require.NoError(t, err) + assert.Equal(t, data[:n], encoded, "re-encoding a decoded field must reproduce the consumed bytes") + }) +} + +// TestField_RoundTrip pins encode→decode symmetry for fields built with NewField. +func TestField_RoundTrip(t *testing.T) { + tests := []struct { + name string + field Field + }{ + {name: "with data", field: NewField(FieldData, []byte("hello"))}, + {name: "empty data", field: NewField(FieldUserPassword, []byte{})}, + {name: "binary data", field: NewField(FieldUserIconID, []byte{0x07, 0xd1})}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + encoded, err := io.ReadAll(&tt.field) + require.NoError(t, err) + + var decoded Field + n, err := decoded.Write(encoded) + require.NoError(t, err) + assert.Equal(t, len(encoded), n) + assert.Equal(t, tt.field.Type, decoded.Type) + assert.Equal(t, tt.field.FieldSize, decoded.FieldSize) + assert.Equal(t, tt.field.Data, decoded.Data) + }) + } +} diff --git a/hotline/file_resume_data_test.go b/hotline/file_resume_data_test.go new file mode 100644 index 0000000..657be76 --- /dev/null +++ b/hotline/file_resume_data_test.go @@ -0,0 +1,87 @@ +package hotline + +import ( + "encoding/binary" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// UnmarshalBinary is covered in decode_test.go. This file covers the constructors, the marshal +// path, encode↔decode symmetry, and ForkType.String(). + +func TestForkType_String(t *testing.T) { + assert.Equal(t, "DATA", ForkTypeDATA.String()) + assert.Equal(t, "INFO", ForkTypeINFO.String()) + assert.Equal(t, "MACR", ForkTypeMACR.String()) +} + +func TestNewForkInfoList(t *testing.T) { + fil := NewForkInfoList([]byte{0x00, 0x00, 0x10, 0x00}) + assert.Equal(t, ForkTypeDATA, fil.Fork) + assert.Equal(t, [4]byte{0x00, 0x00, 0x10, 0x00}, fil.DataSize) + assert.Equal(t, uint32(0x1000), binary.BigEndian.Uint32(fil.DataSize[:])) +} + +func TestNewFileResumeData(t *testing.T) { + frd := NewFileResumeData([]ForkInfoList{*NewForkInfoList([]byte{0, 0, 0, 5})}) + + assert.Equal(t, FormatRFLT, frd.Format) + assert.Equal(t, [2]byte{0, 1}, frd.Version) + assert.Equal(t, [2]byte{0, 1}, frd.ForkCount, "ForkCount low byte tracks the list length") + require.Len(t, frd.ForkInfoList, 1) +} + +func TestFileResumeData_BinaryMarshal(t *testing.T) { + frd := NewFileResumeData([]ForkInfoList{ + *NewForkInfoList([]byte{0, 0, 0, 5}), + *NewForkInfoList([]byte{0, 0, 0, 9}), + }) + + b, err := frd.BinaryMarshal() + require.NoError(t, err) + assert.Len(t, b, resumeDataHeaderLen+2*forkInfoLen) + assert.Equal(t, FormatRFLT[:], b[0:4]) + assert.Equal(t, byte(2), b[41], "ForkCount low byte") +} + +func TestFileResumeData_RoundTrip(t *testing.T) { + tests := []struct { + name string + forks []ForkInfoList + }{ + { + name: "two forks", + forks: []ForkInfoList{*NewForkInfoList([]byte{0, 0, 0, 5}), *NewForkInfoList([]byte{0, 0, 1, 0})}, + }, + { + name: "three forks", + forks: []ForkInfoList{ + *NewForkInfoList([]byte{0, 0, 0, 5}), + *NewForkInfoList([]byte{0, 0, 1, 0}), + *NewForkInfoList([]byte{0, 0, 0, 1}), + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + original := NewFileResumeData(tt.forks) + + b, err := original.BinaryMarshal() + require.NoError(t, err) + + var decoded FileResumeData + require.NoError(t, decoded.UnmarshalBinary(b)) + + // Format/Version/ForkCount and every fork survive the round trip. BinaryMarshal writes + // with LittleEndian and UnmarshalBinary reads with BigEndian, but all fields are + // [n]byte arrays, so the encoding is endian-neutral — this test pins that. + assert.Equal(t, original.Format, decoded.Format) + assert.Equal(t, original.Version, decoded.Version) + assert.Equal(t, original.ForkCount, decoded.ForkCount) + assert.Equal(t, original.ForkInfoList, decoded.ForkInfoList) + }) + } +} diff --git a/hotline/flattened_file_object.go b/hotline/flattened_file_object.go index 9a4fd85..b038b12 100644 --- a/hotline/flattened_file_object.go +++ b/hotline/flattened_file_object.go @@ -12,6 +12,10 @@ import ( // that precedes the variable-length name (and optional comment). const flatFileInfoForkMinLen = 72 +// flatFileInfoForkMaxLen bounds the information fork size declared in an untrusted fork header: +// the fixed-size fields plus the maximum name and comment (both uint16-length-prefixed). +const flatFileInfoForkMaxLen = flatFileInfoForkMinLen + 65535 + 2 + 65535 + type flattenedFileObject struct { FlatFileHeader FlatFileHeader FlatFileInformationForkHeader FlatFileForkHeader @@ -254,6 +258,9 @@ func (ffo *flattenedFileObject) ReadFrom(r io.Reader) (int64, error) { } dataLen := binary.BigEndian.Uint32(ffo.FlatFileInformationForkHeader.DataSize[:]) + if dataLen > flatFileInfoForkMaxLen { + return n, fmt.Errorf("flat file information fork size %d exceeds maximum %d", dataLen, flatFileInfoForkMaxLen) + } ffifBuf := make([]byte, dataLen) if _, err := io.ReadFull(r, ffifBuf); err != nil { return n, err diff --git a/hotline/flattened_file_object_fuzz_test.go b/hotline/flattened_file_object_fuzz_test.go new file mode 100644 index 0000000..574fed3 --- /dev/null +++ b/hotline/flattened_file_object_fuzz_test.go @@ -0,0 +1,95 @@ +package hotline + +import ( + "bytes" + "io" + "testing" + + "github.com/stretchr/testify/require" +) + +// sampleFFOBytes encodes a small flattened file object (FILP header, INFO fork header, +// information fork, DATA fork header) — the layout that ReadFrom parses off the wire. +func sampleFFOBytes(t interface{ Fatal(...any) }) []byte { + ffo := flattenedFileObject{ + FlatFileHeader: FlatFileHeader{ + Format: [4]byte{'F', 'I', 'L', 'P'}, + Version: [2]byte{0, 1}, + ForkCount: [2]byte{0, 2}, + }, + FlatFileInformationFork: NewFlatFileInformationFork("testfile.txt", [8]byte{}, "TEXT", "TTXT"), + FlatFileDataForkHeader: FlatFileForkHeader{ + ForkType: [4]byte{'D', 'A', 'T', 'A'}, + DataSize: [4]byte{0, 0, 0, 5}, + }, + } + b, err := io.ReadAll(&ffo) + if err != nil { + t.Fatal(err) + } + return b +} + +// FuzzFlatFileInformationForkUnmarshal feeds raw untrusted bytes to the information fork +// decoder, which parses the metadata section of client file uploads. +func FuzzFlatFileInformationForkUnmarshal(f *testing.F) { + fork := NewFlatFileInformationFork("testfile.txt", [8]byte{}, "TEXT", "TTXT") + forkBytes, err := io.ReadAll(&fork) + if err != nil { + f.Fatal(err) + } + f.Add(forkBytes) + f.Add(forkBytes[:flatFileInfoForkMinLen]) // fixed-size section only + f.Add([]byte{}) // empty + + f.Fuzz(func(t *testing.T, data []byte) { + var ffif FlatFileInformationFork + if err := ffif.UnmarshalBinary(data); err != nil { + return + } + // A successfully decoded fork must re-encode without panicking. + _, err := io.ReadAll(&ffif) + require.NoError(t, err) + }) +} + +// FuzzFlattenedFileObjectReadFrom feeds raw untrusted bytes to the flattened file object +// parser — the entry point for decoding client file uploads on the transfer port. +func FuzzFlattenedFileObjectReadFrom(f *testing.F) { + sample := sampleFFOBytes(f) + f.Add(sample) + f.Add(sample[:24]) // FILP header only + // Information fork header declaring an absurd size (previously triggered an unbounded allocation). + huge := bytes.Clone(sample) + copy(huge[36:40], []byte{0xff, 0xff, 0xff, 0xff}) + f.Add(huge) + + f.Fuzz(func(t *testing.T, data []byte) { + var ffo flattenedFileObject + _, _ = ffo.ReadFrom(bytes.NewReader(data)) // must not panic or over-allocate + }) +} + +// TestFlattenedFileObject_ReadFrom_RoundTrip pins that ReadFrom can parse what Read encodes. +func TestFlattenedFileObject_ReadFrom_RoundTrip(t *testing.T) { + sample := sampleFFOBytes(t) + + var ffo flattenedFileObject + _, err := ffo.ReadFrom(bytes.NewReader(sample)) + require.NoError(t, err) + + require.Equal(t, [4]byte{'F', 'I', 'L', 'P'}, ffo.FlatFileHeader.Format) + require.Equal(t, []byte("testfile.txt"), ffo.FlatFileInformationFork.Name) + require.Equal(t, ForkType{'D', 'A', 'T', 'A'}, ffo.FlatFileDataForkHeader.ForkType) + require.Equal(t, int64(5), ffo.dataSize()) +} + +// TestFlattenedFileObject_ReadFrom_RejectsOversizedInfoFork pins the untrusted-size guard. +func TestFlattenedFileObject_ReadFrom_RejectsOversizedInfoFork(t *testing.T) { + sample := sampleFFOBytes(t) + copy(sample[36:40], []byte{0xff, 0xff, 0xff, 0xff}) + + var ffo flattenedFileObject + _, err := ffo.ReadFrom(bytes.NewReader(sample)) + require.ErrorContains(t, err, "exceeds maximum") +} diff --git a/hotline/r2_file_store.go b/hotline/r2_file_store.go index 2786b1e..6d4327b 100644 --- a/hotline/r2_file_store.go +++ b/hotline/r2_file_store.go @@ -57,14 +57,14 @@ type s3API interface { // s3Uploader streams a body to R2 as a (multipart, if large) upload. *manager.Uploader satisfies it. type s3Uploader interface { - Upload(ctx context.Context, in *s3.PutObjectInput, opts ...func(*manager.Uploader)) (*manager.UploadOutput, error) + Upload(ctx context.Context, in *s3.PutObjectInput, opts ...func(*manager.Uploader)) (*manager.UploadOutput, error) //nolint:staticcheck // SA1019: transfermanager successor is not yet GA } // NewR2FileStore builds an R2-backed FileStore from a configured S3 client. prefix is an optional // key prefix within the bucket; stagingDir is a local directory used to buffer in-progress // (.incomplete) uploads before they are promoted to R2. func NewR2FileStore(client *s3.Client, bucket, prefix, stagingDir string) *R2FileStore { - return newR2FileStore(client, manager.NewUploader(client), bucket, prefix, stagingDir) + return newR2FileStore(client, manager.NewUploader(client), bucket, prefix, stagingDir) //nolint:staticcheck // SA1019: transfermanager successor is not yet GA } func newR2FileStore(api s3API, up s3Uploader, bucket, prefix, stagingDir string) *R2FileStore { @@ -148,14 +148,14 @@ func (s *R2FileStore) ReadFile(name string) ([]byte, error) { if err != nil { return nil, err } - defer r.Close() + defer func() { _ = r.Close() }() return io.ReadAll(r) } r, err := s.Open(name) if err != nil { return nil, err } - defer r.Close() + defer func() { _ = r.Close() }() return io.ReadAll(r) } @@ -400,7 +400,7 @@ func (s *R2FileStore) promote(oldName, newName string) error { if err != nil { return err // already fs.ErrNotExist-compatible } - defer f.Close() + defer func() { _ = f.Close() }() if _, err := s.uploader.Upload(context.Background(), &s3.PutObjectInput{ Bucket: aws.String(s.bucket), diff --git a/hotline/r2_file_store_test.go b/hotline/r2_file_store_test.go index 5126856..2cff019 100644 --- a/hotline/r2_file_store_test.go +++ b/hotline/r2_file_store_test.go @@ -30,10 +30,25 @@ type fakeS3 struct { mu sync.Mutex objects map[string][]byte modTime map[string]time.Time + + // failOp injects an error the next time the named operation runs (e.g. "GetObject"). + // The entry is consumed on use so callers can target a single call. + failOp map[string]error } func newFakeS3() *fakeS3 { - return &fakeS3{objects: map[string][]byte{}, modTime: map[string]time.Time{}} + return &fakeS3{objects: map[string][]byte{}, modTime: map[string]time.Time{}, failOp: map[string]error{}} +} + +// fail returns and consumes an injected error for op, or nil if none is set. +func (f *fakeS3) fail(op string) error { + f.mu.Lock() + defer f.mu.Unlock() + if err, ok := f.failOp[op]; ok { + delete(f.failOp, op) + return err + } + return nil } func (f *fakeS3) put(key string, data []byte) { @@ -44,6 +59,9 @@ func (f *fakeS3) put(key string, data []byte) { } func (f *fakeS3) HeadObject(_ context.Context, in *s3.HeadObjectInput, _ ...func(*s3.Options)) (*s3.HeadObjectOutput, error) { + if err := f.fail("HeadObject"); err != nil { + return nil, err + } f.mu.Lock() defer f.mu.Unlock() data, ok := f.objects[aws.ToString(in.Key)] @@ -57,6 +75,9 @@ func (f *fakeS3) HeadObject(_ context.Context, in *s3.HeadObjectInput, _ ...func } func (f *fakeS3) GetObject(_ context.Context, in *s3.GetObjectInput, _ ...func(*s3.Options)) (*s3.GetObjectOutput, error) { + if err := f.fail("GetObject"); err != nil { + return nil, err + } f.mu.Lock() defer f.mu.Unlock() data, ok := f.objects[aws.ToString(in.Key)] @@ -70,6 +91,9 @@ func (f *fakeS3) GetObject(_ context.Context, in *s3.GetObjectInput, _ ...func(* } func (f *fakeS3) PutObject(_ context.Context, in *s3.PutObjectInput, _ ...func(*s3.Options)) (*s3.PutObjectOutput, error) { + if err := f.fail("PutObject"); err != nil { + return nil, err + } data, err := io.ReadAll(in.Body) if err != nil { return nil, err @@ -79,7 +103,10 @@ func (f *fakeS3) PutObject(_ context.Context, in *s3.PutObjectInput, _ ...func(* } // Upload satisfies s3Uploader; the fake treats it identically to PutObject. -func (f *fakeS3) Upload(_ context.Context, in *s3.PutObjectInput, _ ...func(*manager.Uploader)) (*manager.UploadOutput, error) { +func (f *fakeS3) Upload(_ context.Context, in *s3.PutObjectInput, _ ...func(*manager.Uploader)) (*manager.UploadOutput, error) { //nolint:staticcheck // SA1019: transfermanager successor is not yet GA + if err := f.fail("Upload"); err != nil { + return nil, err + } data, err := io.ReadAll(in.Body) if err != nil { return nil, err @@ -89,6 +116,9 @@ func (f *fakeS3) Upload(_ context.Context, in *s3.PutObjectInput, _ ...func(*man } func (f *fakeS3) DeleteObject(_ context.Context, in *s3.DeleteObjectInput, _ ...func(*s3.Options)) (*s3.DeleteObjectOutput, error) { + if err := f.fail("DeleteObject"); err != nil { + return nil, err + } f.mu.Lock() defer f.mu.Unlock() delete(f.objects, aws.ToString(in.Key)) @@ -97,6 +127,9 @@ func (f *fakeS3) DeleteObject(_ context.Context, in *s3.DeleteObjectInput, _ ... } func (f *fakeS3) DeleteObjects(_ context.Context, in *s3.DeleteObjectsInput, _ ...func(*s3.Options)) (*s3.DeleteObjectsOutput, error) { + if err := f.fail("DeleteObjects"); err != nil { + return nil, err + } f.mu.Lock() defer f.mu.Unlock() for _, o := range in.Delete.Objects { @@ -107,6 +140,9 @@ func (f *fakeS3) DeleteObjects(_ context.Context, in *s3.DeleteObjectsInput, _ . } func (f *fakeS3) CopyObject(_ context.Context, in *s3.CopyObjectInput, _ ...func(*s3.Options)) (*s3.CopyObjectOutput, error) { + if err := f.fail("CopyObject"); err != nil { + return nil, err + } src, err := decodeCopySource(aws.ToString(in.CopySource)) if err != nil { return nil, err @@ -135,6 +171,9 @@ func decodeCopySource(s string) (string, error) { } func (f *fakeS3) ListObjectsV2(_ context.Context, in *s3.ListObjectsV2Input, _ ...func(*s3.Options)) (*s3.ListObjectsV2Output, error) { + if err := f.fail("ListObjectsV2"); err != nil { + return nil, err + } f.mu.Lock() defer f.mu.Unlock() @@ -395,6 +434,76 @@ func TestR2FileStore_DownloadRoundTrip(t *testing.T) { assert.Equal(t, int64(len(fileData)), ft.bytesSentCounter.Total) } +// TestR2FileStore_ErrorPropagation verifies that S3 API failures surface as errors rather than +// being swallowed or panicking. Errors are injected via fakeS3.failOp. +func TestR2FileStore_ErrorPropagation(t *testing.T) { + injected := errors.New("r2 unavailable") + + t.Run("Open surfaces GetObject error", func(t *testing.T) { + s, fake := newTestR2Store(t) + require.NoError(t, s.WriteFile("/a.txt", []byte("x"), 0644)) + fake.failOp["GetObject"] = injected + + _, err := s.Open("/a.txt") + require.ErrorIs(t, err, injected) + }) + + t.Run("Stat surfaces HeadObject error", func(t *testing.T) { + s, fake := newTestR2Store(t) + require.NoError(t, s.WriteFile("/a.txt", []byte("x"), 0644)) + fake.failOp["HeadObject"] = injected + + _, err := s.Stat("/a.txt") + require.ErrorIs(t, err, injected) + }) + + t.Run("WriteFile surfaces PutObject error", func(t *testing.T) { + s, fake := newTestR2Store(t) + fake.failOp["PutObject"] = injected + + require.ErrorIs(t, s.WriteFile("/a.txt", []byte("x"), 0644), injected) + }) + + t.Run("Remove surfaces DeleteObject error", func(t *testing.T) { + s, fake := newTestR2Store(t) + require.NoError(t, s.WriteFile("/a.txt", []byte("x"), 0644)) + fake.failOp["DeleteObject"] = injected + + require.ErrorIs(t, s.Remove("/a.txt"), injected) + }) + + t.Run("RemoveAll surfaces ListObjectsV2 error", func(t *testing.T) { + s, fake := newTestR2Store(t) + require.NoError(t, s.WriteFile("/dir/a.txt", []byte("x"), 0644)) + fake.failOp["ListObjectsV2"] = injected + + require.ErrorIs(t, s.RemoveAll("/dir"), injected) + }) + + t.Run("Rename of existing object surfaces CopyObject error", func(t *testing.T) { + s, fake := newTestR2Store(t) + require.NoError(t, s.WriteFile("/a.txt", []byte("x"), 0644)) + fake.failOp["CopyObject"] = injected + + require.ErrorIs(t, s.Rename("/a.txt", "/b.txt"), injected) + }) + + t.Run("upload commit surfaces Upload error and leaves staged file", func(t *testing.T) { + s, fake := newTestR2Store(t) + incomplete := "/upload.txt" + IncompleteFileSuffix + require.NoError(t, s.WriteFile(incomplete, []byte("partial"), 0644)) + fake.failOp["Upload"] = injected + + // promote() streams the staged file up on the terminal rename; a failure must surface. + err := s.Rename(incomplete, "/upload.txt") + require.ErrorIs(t, err, injected) + + // The staged file is retained so the transfer can resume. + _, statErr := s.Stat(incomplete) + require.NoError(t, statErr) + }) +} + // TestR2FileStore_Integration round-trips against a real R2 bucket. It is skipped unless R2_BUCKET // and credentials are present in the environment. func TestR2FileStore_Integration(t *testing.T) { diff --git a/hotline/server.go b/hotline/server.go index 4bf48c6..2547c38 100644 --- a/hotline/server.go +++ b/hotline/server.go @@ -74,6 +74,12 @@ type Server struct { TLSConfig *tls.Config TLSPort int + // connRateLimit and connRateBurst bound how frequently a single IP may connect. They default to + // the production values (perIPRateLimit, 1) and are overridable via WithConnectionRateLimit, + // which tests set to rate.Inf to avoid the 2s-per-connection throttle. + connRateLimit rate.Limit + connRateBurst int + shutdownInit sync.Once // lazily creates shutdownCh so Shutdown works on test-constructed Servers shutdownOnce sync.Once // guards close(shutdownCh) shutdownCh chan struct{} // closed by Shutdown to stop ListenAndServe @@ -203,6 +209,16 @@ func WithTLS(tlsConfig *tls.Config, port int) func(s *Server) { } } +// WithConnectionRateLimit overrides the per-IP connection rate limit. limit is the sustained rate +// (connections per second) and burst the bucket size. Pass rate.Inf to disable throttling, e.g. in +// tests that open several connections in quick succession. +func WithConnectionRateLimit(limit rate.Limit, burst int) func(s *Server) { + return func(s *Server) { + s.connRateLimit = limit + s.connRateBurst = burst + } +} + type ServerConfig struct { } @@ -216,6 +232,8 @@ func NewServer(options ...Option) (*Server, error) { FileTransferMgr: NewMemFileTransferMgr(), Stats: NewStats(), TrackerRegistrar: NewRealTrackerRegistrar(), + connRateLimit: perIPRateLimit, + connRateBurst: 1, } for _, opt := range options { @@ -430,7 +448,7 @@ func (s *Server) Serve(ctx context.Context, ln net.Listener) error { s.rateLimitersMu.Lock() entry, ok := s.rateLimiters[ipAddr] if !ok { - entry = &rateLimiterEntry{limiter: rate.NewLimiter(perIPRateLimit, 1)} + entry = &rateLimiterEntry{limiter: rate.NewLimiter(s.connRateLimit, s.connRateBurst)} s.rateLimiters[ipAddr] = entry } entry.lastSeen = time.Now() @@ -741,8 +759,16 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser c.Send(NewTransaction(TranShowAgreement, c.ID, NewField(FieldNoServerAgreement, []byte{1}))) } } else { - _, _ = c.Server.Agreement.Seek(0, 0) - data, _ := io.ReadAll(c.Server.Agreement) + // Prefer a snapshot read: the Agreement is shared across all connections, so the stateful + // Seek+ReadAll below races when clients log in concurrently. Implementations that expose + // AgreementBytes return a private copy safely. + var data []byte + if ab, ok := c.Server.Agreement.(interface{ AgreementBytes() []byte }); ok { + data = ab.AgreementBytes() + } else { + _, _ = c.Server.Agreement.Seek(0, 0) + data, _ = io.ReadAll(c.Server.Agreement) + } c.Send(NewTransaction(TranShowAgreement, c.ID, NewField(FieldData, data))) } diff --git a/hotline/server_blackbox_test.go b/hotline/server_blackbox_test.go index 7033703..bfd9e9c 100644 --- a/hotline/server_blackbox_test.go +++ b/hotline/server_blackbox_test.go @@ -16,21 +16,31 @@ func NewTestLogger() *slog.Logger { return slog.New(slog.NewTextHandler(os.Stdout, nil)) } -// assertTransferBytesEqual takes a string with a hexdump in the same format that `hexdump -C` produces and compares with -// a hexdump for the bytes in got, after stripping the create/modify timestamps. -// I don't love this, but as git does not preserve file create/modify timestamps, we either need to fully mock the -// filesystem interactions or work around in this way. -// TODO: figure out a better solution +// flatFileTimestampOffset is the byte offset of the information fork's CreateDate within a +// flattened file object stream: FlatFileHeader (24) + INFO FlatFileForkHeader (16) + the info +// fork's fixed fields up to CreateDate (Platform+TypeSignature+CreatorSignature+Flags+ +// PlatformFlags = 20, then RSVD = 32). CreateDate and ModifyDate are 8 bytes each and adjacent. +const ( + flatFileTimestampOffset = 24 + 16 + 20 + 32 + flatFileTimestampLen = 16 // CreateDate (8) + ModifyDate (8) +) + +// assertTransferBytesEqual takes a string with a hexdump in the same format that `hexdump -C` +// produces and compares with a hexdump for the bytes in got, after zeroing the info fork's +// create/modify timestamps. Git does not preserve file create/modify times, so those bytes vary +// between checkouts; the offset is derived structurally from the flattened file object layout +// rather than hardcoded (see flatFileTimestampOffset). func assertTransferBytesEqual(t *testing.T, wantHexDump string, got []byte) bool { if wantHexDump == "" { return true } - clean := slices.Concat( - got[:92], - make([]byte, 16), - got[108:], - ) + clean := slices.Clone(got) + if len(clean) >= flatFileTimestampOffset+flatFileTimestampLen { + for i := flatFileTimestampOffset; i < flatFileTimestampOffset+flatFileTimestampLen; i++ { + clean[i] = 0 + } + } return assert.Equal(t, wantHexDump, hex.Dump(clean)) } diff --git a/hotline/transaction.go b/hotline/transaction.go index 037e939..b1d5279 100644 --- a/hotline/transaction.go +++ b/hotline/transaction.go @@ -169,6 +169,12 @@ func (t *Transaction) Write(p []byte) (n int, err error) { totalSize := binary.BigEndian.Uint32(p[12:16]) tranLen := int(20 + totalSize) + // The size fields are untrusted network input; reject sizes that fall outside the buffer + // rather than letting the field slice below panic. + if tranLen < 22 || tranLen > len(p) { + return 0, errors.New("invalid transaction size") + } + paramCount := binary.BigEndian.Uint16(p[20:22]) t.Flags = p[0] @@ -213,13 +219,14 @@ func transactionScanner(data []byte, _ bool) (advance int, token []byte, err err totalSize := binary.BigEndian.Uint32(data[12:16]) - // tranLen represents the length of bytes that are part of the transaction - tranLen := int(tranHeaderLen + totalSize) - if tranLen > len(data) { + // tranLen represents the length of bytes that are part of the transaction. Compute it in int64 + // so a near-max totalSize can't overflow (uint32 addition would wrap and yield a short length). + tranLen := int64(tranHeaderLen) + int64(totalSize) + if tranLen > int64(len(data)) { return 0, nil, nil } - return tranLen, data[0:tranLen], nil + return int(tranLen), data[0:tranLen], nil } const minFieldLen = 4 diff --git a/hotline/transaction_fuzz_test.go b/hotline/transaction_fuzz_test.go new file mode 100644 index 0000000..60e7fff --- /dev/null +++ b/hotline/transaction_fuzz_test.go @@ -0,0 +1,149 @@ +package hotline + +import ( + "bytes" + "io" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// sampleTransactionBytes is a valid TranChatSend transaction with one FieldData field ("hai"), +// lifted from TestTransaction_Write. +var sampleTransactionBytes = []byte{ + 0x00, 0x00, 0x00, 0x69, 0x00, 0x00, 0x15, 0x72, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x09, + 0x00, 0x00, 0x00, 0x09, 0x00, 0x01, 0x00, 0x65, + 0x00, 0x03, 0x68, 0x61, 0x69, +} + +// FuzzTransactionScanner verifies that the split func used to frame transactions from the +// network never panics, never advances past its input, and only produces tokens that the +// Transaction decoder can consume without panicking. +func FuzzTransactionScanner(f *testing.F) { + f.Add(sampleTransactionBytes) + f.Add(sampleTransactionBytes[:16]) // header only, no fields + f.Add([]byte{0xff, 0xff, 0xff, 0xff}) // shorter than the size field + f.Add(bytes.Repeat([]byte{0xff}, 64)) // absurd declared size + f.Add(append(sampleTransactionBytes, 0xde, 0xad, 0xbe)) // trailing partial transaction + + f.Fuzz(func(t *testing.T, data []byte) { + advance, token, err := transactionScanner(data, false) + if err != nil { + return + } + if advance == 0 { + assert.Nil(t, token, "no advance must produce no token") + return + } + require.LessOrEqual(t, advance, len(data), "scanner advanced past its input") + require.Len(t, token, advance, "token length must match advance") + require.GreaterOrEqual(t, advance, tranHeaderLen, "token cannot be smaller than a transaction header") + + // A framed token must be decodable without panicking (errors are fine). + _, _ = (&Transaction{}).Write(token) + }) +} + +// FuzzTransactionWrite feeds raw untrusted bytes to the Transaction decoder. The size and +// param-count fields come straight off the network, so decoding must error rather than panic. +func FuzzTransactionWrite(f *testing.F) { + f.Add(sampleTransactionBytes) + f.Add(sampleTransactionBytes[:22]) + // Declared total size smaller than the minimum (previously panicked with tranLen < 22). + f.Add([]byte{ + 0x00, 0x00, 0x00, 0x69, 0x00, 0x00, 0x15, 0x72, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + }) + // Declared total size larger than the buffer (previously panicked with tranLen > len(p)). + f.Add([]byte{ + 0x00, 0x00, 0x00, 0x69, 0x00, 0x00, 0x15, 0x72, + 0x00, 0x00, 0x00, 0x00, 0xff, 0xff, 0xff, 0xff, + 0x00, 0x00, 0x00, 0x09, 0x00, 0x01, + }) + + f.Fuzz(func(t *testing.T, data []byte) { + var tran Transaction + if _, err := tran.Write(data); err != nil { + return + } + + // A successfully decoded transaction must re-encode and decode back to the same value. + encoded, err := io.ReadAll(&tran) + require.NoError(t, err) + + var reDecoded Transaction + _, err = reDecoded.Write(encoded) + require.NoError(t, err, "re-encoded transaction failed to decode") + assert.Equal(t, tran.Type, reDecoded.Type) + assert.Equal(t, tran.Fields, reDecoded.Fields) + }) +} + +// TestTransaction_RoundTrip pins encode→decode symmetry for representative transactions. +func TestTransaction_RoundTrip(t *testing.T) { + tests := []struct { + name string + tran Transaction + }{ + { + name: "no fields", + tran: Transaction{ + Type: TranKeepAlive, + ID: [4]byte{0, 0, 0, 1}, + }, + }, + { + name: "single field", + tran: Transaction{ + Type: TranChatSend, + ID: [4]byte{0, 0, 0, 2}, + Fields: []Field{NewField(FieldData, []byte("hello world"))}, + }, + }, + { + name: "multiple fields including empty data", + tran: Transaction{ + IsReply: 1, + ErrorCode: [4]byte{0, 0, 0, 1}, + Type: TranLogin, + ID: [4]byte{0, 0, 0, 3}, + Fields: []Field{ + NewField(FieldUserLogin, []byte("guest")), + NewField(FieldUserPassword, []byte{}), + NewField(FieldUserIconID, []byte{0x07, 0xd1}), + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + encoded, err := io.ReadAll(&tt.tran) + require.NoError(t, err) + + var decoded Transaction + n, err := decoded.Write(encoded) + require.NoError(t, err) + assert.Equal(t, len(encoded), n) + + assert.Equal(t, tt.tran.Flags, decoded.Flags) + assert.Equal(t, tt.tran.IsReply, decoded.IsReply) + assert.Equal(t, tt.tran.Type, decoded.Type) + assert.Equal(t, tt.tran.ID, decoded.ID) + assert.Equal(t, tt.tran.ErrorCode, decoded.ErrorCode) + if len(tt.tran.Fields) == 0 { + assert.Empty(t, decoded.Fields) + } else { + assert.Equal(t, tt.tran.Fields, decoded.Fields) + } + + // Encoding the decoded transaction must reproduce the original bytes. + reEncoded, err := io.ReadAll(&decoded) + require.NoError(t, err) + assert.Equal(t, encoded, reEncoded) + }) + } +} |