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 | |
| 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.
| -rw-r--r-- | .github/workflows/golangci-lint.yml | 9 | ||||
| -rw-r--r-- | .github/workflows/test.yml | 34 | ||||
| -rw-r--r-- | Makefile | 22 | ||||
| -rw-r--r-- | cmd/mobius-hotline-server/main_test.go | 80 | ||||
| -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 | ||||
| -rw-r--r-- | internal/mobius/account_manager.go | 12 | ||||
| -rw-r--r-- | internal/mobius/agreement.go | 14 | ||||
| -rw-r--r-- | internal/mobius/config_test.go | 14 | ||||
| -rw-r--r-- | internal/mobius/e2e_test.go | 439 | ||||
| -rw-r--r-- | internal/mobius/integration_test.go | 429 | ||||
| -rw-r--r-- | internal/mobius/reload_test.go | 32 |
22 files changed, 1668 insertions, 67 deletions
diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index 6f2d802..d788728 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -1,7 +1,8 @@ name: golangci-lint on: + push: + branches: [master] pull_request: - types: [opened, reopened] permissions: contents: read @@ -18,8 +19,8 @@ jobs: - uses: actions/checkout@v4 - uses: actions/setup-go@v5 with: - go-version: stable + go-version-file: go.mod - name: golangci-lint - uses: golangci/golangci-lint-action@v6 + uses: golangci/golangci-lint-action@v8 with: - version: v1.58
\ No newline at end of file + version: v2.6
\ No newline at end of file diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml new file mode 100644 index 0000000..3d4e155 --- /dev/null +++ b/.github/workflows/test.yml @@ -0,0 +1,34 @@ +name: test + +on: + push: + branches: [master] + pull_request: + +permissions: + contents: read + +jobs: + test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + - name: Build + run: go build ./... + - name: Vet + run: go vet ./... + - name: Test + run: go test ./... -race -shuffle=on -coverprofile=coverage.out -covermode=atomic + - name: Coverage summary + run: | + echo '### Coverage' >> "$GITHUB_STEP_SUMMARY" + echo '```' >> "$GITHUB_STEP_SUMMARY" + go tool cover -func=coverage.out | tail -1 >> "$GITHUB_STEP_SUMMARY" + echo '```' >> "$GITHUB_STEP_SUMMARY" + - uses: actions/upload-artifact@v4 + with: + name: coverage + path: coverage.out @@ -1,2 +1,24 @@ server: go build -ldflags "-X main.version=$$(git describe --exact-match --tags || echo "dev" ) -X main.commit=$$(git rev-parse --short HEAD)" -o mobius-hotline-server cmd/mobius-hotline-server/main.go + +.PHONY: test +test: + go test ./... -race -shuffle=on + +.PHONY: cover +cover: + go test ./... -race -shuffle=on -coverprofile=coverage.out -covermode=atomic + go tool cover -html=coverage.out + +.PHONY: lint +lint: + golangci-lint run + +# Run each fuzz target for FUZZTIME (default 30s). Plain `go test` already replays +# committed seed corpora; this exercises new random inputs. +.PHONY: fuzz +fuzz: + @for target in $$(go test ./hotline -list 'Fuzz.*' | grep '^Fuzz'); do \ + echo "fuzzing $$target"; \ + go test ./hotline -run '^$$' -fuzz "^$$target$$" -fuzztime $${FUZZTIME:-30s} || exit 1; \ + done diff --git a/cmd/mobius-hotline-server/main_test.go b/cmd/mobius-hotline-server/main_test.go index a527d8b..827397a 100644 --- a/cmd/mobius-hotline-server/main_test.go +++ b/cmd/mobius-hotline-server/main_test.go @@ -1,6 +1,7 @@ package main import ( + "context" "os" "path" "testing" @@ -179,3 +180,82 @@ func TestFindConfigPath(t *testing.T) { assert.Equal(t, "config", result) }) } + +// r2EnvVars are every R2_* variable newR2FileStore reads. Each case clears all of them and sets +// only what it needs, so ambient credentials in the test environment can't leak in. +var r2EnvVars = []string{ + "R2_BUCKET", "R2_ACCESS_KEY_ID", "R2_SECRET_ACCESS_KEY", + "R2_ENDPOINT", "R2_ACCOUNT_ID", "R2_PREFIX", "R2_STAGING_DIR", +} + +func TestNewR2FileStore_EnvValidation(t *testing.T) { + tests := []struct { + name string + env map[string]string + wantErr string // substring; empty means the store must construct successfully + }{ + { + name: "missing bucket and credentials", + env: map[string]string{}, + wantErr: "R2_BUCKET, R2_ACCESS_KEY_ID, and R2_SECRET_ACCESS_KEY must be set", + }, + { + name: "missing secret key", + env: map[string]string{ + "R2_BUCKET": "b", + "R2_ACCESS_KEY_ID": "ak", + }, + wantErr: "R2_BUCKET, R2_ACCESS_KEY_ID, and R2_SECRET_ACCESS_KEY must be set", + }, + { + name: "credentials set but no endpoint or account id", + env: map[string]string{ + "R2_BUCKET": "b", + "R2_ACCESS_KEY_ID": "ak", + "R2_SECRET_ACCESS_KEY": "sk", + }, + wantErr: "either R2_ENDPOINT or R2_ACCOUNT_ID must be set", + }, + { + name: "account id derives the endpoint", + env: map[string]string{ + "R2_BUCKET": "b", + "R2_ACCESS_KEY_ID": "ak", + "R2_SECRET_ACCESS_KEY": "sk", + "R2_ACCOUNT_ID": "acct123", + }, + }, + { + name: "explicit endpoint", + env: map[string]string{ + "R2_BUCKET": "b", + "R2_ACCESS_KEY_ID": "ak", + "R2_SECRET_ACCESS_KEY": "sk", + "R2_ENDPOINT": "https://example.com", + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + for _, k := range r2EnvVars { + t.Setenv(k, "") + } + // Keep staging off the shared temp dir even on the success paths. + t.Setenv("R2_STAGING_DIR", t.TempDir()) + for k, v := range tt.env { + t.Setenv(k, v) + } + + store, err := newR2FileStore(context.Background()) + if tt.wantErr != "" { + require.Error(t, err) + assert.Contains(t, err.Error(), tt.wantErr) + assert.Nil(t, store) + return + } + require.NoError(t, err) + assert.NotNil(t, store) + }) + } +} 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) + }) + } +} diff --git a/internal/mobius/account_manager.go b/internal/mobius/account_manager.go index 1265d9c..44c054a 100644 --- a/internal/mobius/account_manager.go +++ b/internal/mobius/account_manager.go @@ -13,18 +13,6 @@ import ( "gopkg.in/yaml.v3" ) -// loadFromYAMLFile loads data from a YAML file into the provided data structure. -func loadFromYAMLFile(path string, data interface{}) error { - fh, err := os.Open(path) - if err != nil { - return err - } - defer func() { _ = fh.Close() }() - - decoder := yaml.NewDecoder(fh) - return decoder.Decode(data) -} - // YAMLAccountManager implements AccountManager interface using YAML files for persistence. // It maintains an in-memory cache of accounts and synchronizes with YAML files on disk. type YAMLAccountManager struct { diff --git a/internal/mobius/agreement.go b/internal/mobius/agreement.go index d9cf2d4..0f652e5 100644 --- a/internal/mobius/agreement.go +++ b/internal/mobius/agreement.go @@ -63,7 +63,21 @@ func (a *Agreement) Read(p []byte) (int, error) { } func (a *Agreement) Seek(offset int64, _ int) (int64, error) { + a.mu.Lock() + defer a.mu.Unlock() + a.readOffset = int(offset) return 0, nil } + +// AgreementBytes returns a private copy of the agreement text. Unlike Seek+Read it touches no +// shared read offset, so concurrent logins can each obtain the agreement without racing. +func (a *Agreement) AgreementBytes() []byte { + a.mu.RLock() + defer a.mu.RUnlock() + + out := make([]byte, len(a.data)) + copy(out, a.data) + return out +} diff --git a/internal/mobius/config_test.go b/internal/mobius/config_test.go index bed739d..f090b5c 100644 --- a/internal/mobius/config_test.go +++ b/internal/mobius/config_test.go @@ -9,11 +9,7 @@ import ( func TestLoadConfig_InvalidBannerFileExtension(t *testing.T) { // Create a temporary directory for test files - tmpDir, err := os.MkdirTemp("", "mobius-config-test") - if err != nil { - t.Fatalf("Failed to create temp dir: %v", err) - } - defer os.RemoveAll(tmpDir) + tmpDir := t.TempDir() // Create a test config file with an invalid banner file extension configContent := ` @@ -28,7 +24,7 @@ FileRoot: "files" } // Attempt to load the config - _, err = LoadConfig(configPath) + _, err := LoadConfig(configPath) // Verify that we get the improved error message if err == nil { @@ -94,11 +90,7 @@ func TestLoadConfig_ValidBannerFileExtensions(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { // Create a temporary directory for test files - tmpDir, err := os.MkdirTemp("", "mobius-config-test") - if err != nil { - t.Fatalf("Failed to create temp dir: %v", err) - } - defer os.RemoveAll(tmpDir) + tmpDir := t.TempDir() // Create files subdirectory filesDir := filepath.Join(tmpDir, "files") diff --git a/internal/mobius/e2e_test.go b/internal/mobius/e2e_test.go new file mode 100644 index 0000000..ccd56f3 --- /dev/null +++ b/internal/mobius/e2e_test.go @@ -0,0 +1,439 @@ +package mobius + +// End-to-end protocol regression tests. Each top-level test starts a fresh in-process server (see +// integration_test.go for the harness) and drives it with the real hotline.Client over TCP. + +import ( + "bytes" + "encoding/binary" + "io" + "net" + "os" + "path/filepath" + "testing" + "time" + + "github.com/jhalter/mobius/hotline" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func mustReadAll(r io.Reader) []byte { + b, err := io.ReadAll(r) + if err != nil { + panic(err) + } + return b +} + +// readUploadedFile polls for an uploaded file to appear under the server's Files tree (the upload +// completes asynchronously on the transfer connection) and returns its data-fork content. +func readUploadedFile(t *testing.T, s *e2eServer, name string) []byte { + t.Helper() + path := filepath.Join(s.cfgDir, "Files", name) + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + if data, err := os.ReadFile(path); err == nil { + return data + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("uploaded file %s never appeared", path) + return nil +} + +func newReq(tranType [2]byte, fields ...hotline.Field) hotline.Transaction { + return hotline.NewTransaction(tranType, [2]byte{0, 0}, fields...) +} + +func TestE2E_Handshake(t *testing.T) { + s := startE2EServer(t) + + t.Run("valid handshake succeeds", func(t *testing.T) { + c := hotline.NewClient("probe", NewTestLogger()) + conn, err := net.DialTimeout("tcp", s.addr, 5*time.Second) + require.NoError(t, err) + defer func() { _ = conn.Close() }() + c.Connection = conn + require.NoError(t, c.Handshake()) + }) + + t.Run("garbage handshake is rejected", func(t *testing.T) { + conn, err := net.DialTimeout("tcp", s.addr, 5*time.Second) + require.NoError(t, err) + defer func() { _ = conn.Close() }() + + _, err = conn.Write([]byte("not a valid hotline handshake!!!")) + require.NoError(t, err) + + // The server closes the connection on a bad handshake; the read returns EOF. + require.NoError(t, conn.SetReadDeadline(time.Now().Add(3*time.Second))) + buf := make([]byte, 8) + _, err = conn.Read(buf) + require.Error(t, err) + }) +} + +func TestE2E_Login(t *testing.T) { + s := startE2EServer(t) + + t.Run("guest login and agreement flow", func(t *testing.T) { + c := connectE2E(t, s.addr, "guest", "") + // A successful, fully-established session answers a request. The user name list is not + // access-gated, so a non-error reply proves login (and the agreement auto-reply) completed. + reply := c.roundTrip(newReq(hotline.TranGetUserNameList)) + assert.False(t, isError(reply), "an established session should answer TranGetUserNameList") + }) + + t.Run("wrong password is rejected", func(t *testing.T) { + c := hotline.NewClient("baduser", NewTestLogger()) + require.NoError(t, c.Connect(s.addr, e2eAdminLogin, "wrong-password")) + // The server replies to the login with an error, then closes the connection. + require.NoError(t, c.Connection.SetReadDeadline(time.Now().Add(3*time.Second))) + buf := make([]byte, 512) + n, err := c.Connection.Read(buf) + if err == nil { + // If a reply was delivered it must be an error reply. + var reply hotline.Transaction + _, werr := reply.Write(buf[:n]) + if werr == nil { + assert.True(t, isError(reply), "login with wrong password should return an error reply") + } + } + _ = c.Disconnect() + }) +} + +func TestE2E_Chat(t *testing.T) { + s := startE2EServer(t) + + alice := connectE2E(t, s.addr, "guest", "") + bob := connectE2E(t, s.addr, "guest", "") + + // A completed round trip proves a session is published to the client manager (publication + // happens before the server answers requests), so after these both peers are guaranteed to be + // in the chat broadcast set. Waiting for a login notification instead would be order-dependent: + // alice only hears about bob if bob logs in after alice is established. + require.False(t, isError(alice.roundTrip(newReq(hotline.TranGetUserNameList)))) + require.False(t, isError(bob.roundTrip(newReq(hotline.TranGetUserNameList)))) + + require.NoError(t, alice.client.Send(newReq(hotline.TranChatSend, + hotline.NewField(hotline.FieldData, []byte("hello from alice")), + ))) + + msg := bob.waitFor(hotline.TranChatMsg) + got := msg.GetField(hotline.FieldData) + require.NotNil(t, got) + assert.Contains(t, string(got.Data), "hello from alice") +} + +// findUserID picks the user with the given display name out of a TranGetUserNameList reply. +func findUserID(t *testing.T, reply hotline.Transaction, name string) [2]byte { + t.Helper() + for _, f := range reply.Fields { + if f.Type != hotline.FieldUsernameWithInfo { + continue + } + var u hotline.User + _, err := u.Write(f.Data) + require.NoError(t, err) + if u.Name == name { + return u.ID + } + } + t.Fatalf("user %q not in user list", name) + return [2]byte{} +} + +func TestE2E_PrivateChat(t *testing.T) { + s := startE2EServer(t) + + alice := connectE2ENamed(t, s.addr, "alice", "guest", "") + bob := connectE2ENamed(t, s.addr, "bob", "guest", "") + + // A completed round trip proves a session is published to the client manager, so bob is + // guaranteed to be in alice's user list below. + require.False(t, isError(bob.roundTrip(newReq(hotline.TranGetUserNameList)))) + list := alice.roundTrip(newReq(hotline.TranGetUserNameList)) + require.False(t, isError(list)) + bobID := findUserID(t, list, "bob") + + // Alice opens a private chat with bob; the reply carries the new chat ID. + invite := alice.roundTrip(newReq(hotline.TranInviteNewChat, + hotline.NewField(hotline.FieldUserID, bobID[:]), + )) + require.False(t, isError(invite)) + chatIDField := invite.GetField(hotline.FieldChatID) + require.NotNil(t, chatIDField) + chatID := chatIDField.Data + + // Bob receives the invitation, carrying the same chat ID. + inv := bob.waitFor(hotline.TranInviteToChat) + require.Equal(t, chatID, inv.GetField(hotline.FieldChatID).Data) + + // Bob joins; the join reply lists alice as an existing member. + join := bob.roundTrip(newReq(hotline.TranJoinChat, + hotline.NewField(hotline.FieldChatID, chatID), + )) + require.False(t, isError(join)) + var members []string + for _, f := range join.Fields { + if f.Type == hotline.FieldUsernameWithInfo { + var u hotline.User + _, err := u.Write(f.Data) + require.NoError(t, err) + members = append(members, u.Name) + } + } + assert.Contains(t, members, "alice", "join reply should list the chat's existing members") + + // Alice is notified that bob joined. + joined := alice.waitFor(hotline.TranNotifyChatChangeUser) + assert.Equal(t, chatID, joined.GetField(hotline.FieldChatID).Data) + + // Alice sets the subject; bob is notified. + require.NoError(t, alice.client.Send(newReq(hotline.TranSetChatSubject, + hotline.NewField(hotline.FieldChatID, chatID), + hotline.NewField(hotline.FieldChatSubject, []byte("secret plans")), + ))) + subj := bob.waitFor(hotline.TranNotifyChatSubject) + assert.Equal(t, "secret plans", string(subj.GetField(hotline.FieldChatSubject).Data)) + + // A message sent with the chat ID goes to the room's members (including the sender's echo), + // tagged with the chat ID. + require.NoError(t, alice.client.Send(newReq(hotline.TranChatSend, + hotline.NewField(hotline.FieldChatID, chatID), + hotline.NewField(hotline.FieldData, []byte("psst")), + ))) + msg := bob.waitFor(hotline.TranChatMsg) + assert.Equal(t, chatID, msg.GetField(hotline.FieldChatID).Data) + assert.Contains(t, string(msg.GetField(hotline.FieldData).Data), "psst") + echo := alice.waitFor(hotline.TranChatMsg) + assert.Contains(t, string(echo.GetField(hotline.FieldData).Data), "psst") + + // Bob leaves; alice is notified. + require.NoError(t, bob.client.Send(newReq(hotline.TranLeaveChat, + hotline.NewField(hotline.FieldChatID, chatID), + ))) + left := alice.waitFor(hotline.TranNotifyChatDeleteUser) + assert.Equal(t, chatID, left.GetField(hotline.FieldChatID).Data) + + // A declined invitation is announced to the chat's members. + invite2 := alice.roundTrip(newReq(hotline.TranInviteNewChat, + hotline.NewField(hotline.FieldUserID, bobID[:]), + )) + require.False(t, isError(invite2)) + inv2 := bob.waitFor(hotline.TranInviteToChat) + require.NoError(t, bob.client.Send(newReq(hotline.TranRejectChatInvite, + hotline.NewField(hotline.FieldChatID, inv2.GetField(hotline.FieldChatID).Data), + ))) + decline := alice.waitFor(hotline.TranChatMsg) + assert.Contains(t, string(decline.GetField(hotline.FieldData).Data), "bob declined invitation to chat") +} + +func TestE2E_MessageBoard(t *testing.T) { + s := startE2EServer(t) + c := connectE2E(t, s.addr, "guest", "") + + reply := c.roundTrip(newReq(hotline.TranGetMsgs)) + require.False(t, isError(reply)) + data := reply.GetField(hotline.FieldData) + require.NotNil(t, data) + assert.Contains(t, string(data.Data), "Test News Post", "message board should return the fixture content") +} + +func TestE2E_ThreadedNews(t *testing.T) { + s := startE2EServer(t) + c := connectE2E(t, s.addr, "guest", "") + + reply := c.roundTrip(newReq(hotline.TranGetNewsCatNameList, + hotline.NewField(hotline.FieldNewsPath, []byte{}), + )) + require.False(t, isError(reply)) + // The fixture ThreadedNews.yaml defines a "TestBundle" category at the root. + var found bool + for _, f := range reply.Fields { + if bytes.Contains(f.Data, []byte("TestBundle")) { + found = true + } + } + assert.True(t, found, "root news category list should include the fixture bundle") +} + +func TestE2E_FileList(t *testing.T) { + s := startE2EServer(t) + c := connectE2E(t, s.addr, "guest", "") + + reply := c.roundTrip(newReq(hotline.TranGetFileNameList, + hotline.NewField(hotline.FieldFilePath, []byte{}), + )) + require.False(t, isError(reply)) + + var found bool + for _, f := range reply.Fields { + if f.Type == hotline.FieldFileNameWithInfo && bytes.Contains(f.Data, []byte("testfile.txt")) { + found = true + } + } + assert.True(t, found, "file list should include the fixture testfile.txt") +} + +func TestE2E_FileDownload(t *testing.T) { + s := startE2EServer(t) + c := connectE2E(t, s.addr, "guest", "") + + reply := c.roundTrip(newReq(hotline.TranDownloadFile, + hotline.NewField(hotline.FieldFileName, []byte("testfile.txt")), + hotline.NewField(hotline.FieldFilePath, []byte{}), + )) + require.False(t, isError(reply)) + + refField := reply.GetField(hotline.FieldRefNum) + require.NotNil(t, refField, "download reply must carry a reference number") + var refNum [4]byte + copy(refNum[:], refField.Data) + + data := downloadOverTransferPort(t, s.xferAddr, refNum) + assert.Contains(t, string(data), "Hello, I'm a test file!", "downloaded payload should contain the fixture data fork") +} + +func TestE2E_FileUpload(t *testing.T) { + s := startE2EServer(t) + // Upload to the file root requires AccessUploadAnywhere, which the admin account has and guest + // does not. + c := connectE2E(t, s.addr, e2eAdminLogin, e2eAdminPass) + + content := []byte("uploaded by the e2e suite") + payload := buildUploadFFO("uploaded.txt", content) + + reply := c.roundTrip(newReq(hotline.TranUploadFile, + hotline.NewField(hotline.FieldFileName, []byte("uploaded.txt")), + hotline.NewField(hotline.FieldFilePath, []byte{}), + hotline.NewField(hotline.FieldTransferSize, u32(uint32(len(payload)))), + )) + require.False(t, isError(reply)) + + refField := reply.GetField(hotline.FieldRefNum) + require.NotNil(t, refField) + var refNum [4]byte + copy(refNum[:], refField.Data) + + uploadOverTransferPort(t, s.xferAddr, refNum, payload) + + // The uploaded file should land under the server's Files tree with the expected content. + uploaded := readUploadedFile(t, s, "uploaded.txt") + assert.Equal(t, content, uploaded) +} + +func TestE2E_AccountAdmin(t *testing.T) { + s := startE2EServer(t) + + t.Run("admin can create, read, and delete a user", func(t *testing.T) { + admin := connectE2E(t, s.addr, e2eAdminLogin, e2eAdminPass) + + create := admin.roundTrip(newReq(hotline.TranNewUser, + hotline.NewField(hotline.FieldUserLogin, hotline.EncodeString([]byte("newbie"))), + hotline.NewField(hotline.FieldUserName, []byte("New Bie")), + hotline.NewField(hotline.FieldUserPassword, hotline.EncodeString([]byte("pw"))), + hotline.NewField(hotline.FieldUserAccess, make([]byte, 8)), + )) + require.False(t, isError(create), "admin should be allowed to create a user") + + get := admin.roundTrip(newReq(hotline.TranGetUser, + hotline.NewField(hotline.FieldUserLogin, []byte("newbie")), + )) + require.False(t, isError(get)) + name := get.GetField(hotline.FieldUserName) + require.NotNil(t, name) + assert.Equal(t, "New Bie", string(name.Data)) + + del := admin.roundTrip(newReq(hotline.TranDeleteUser, + hotline.NewField(hotline.FieldUserLogin, hotline.EncodeString([]byte("newbie"))), + )) + assert.False(t, isError(del)) + }) + + t.Run("guest cannot create a user", func(t *testing.T) { + guest := connectE2E(t, s.addr, "guest", "") + reply := guest.roundTrip(newReq(hotline.TranNewUser, + hotline.NewField(hotline.FieldUserLogin, hotline.EncodeString([]byte("sneaky"))), + hotline.NewField(hotline.FieldUserName, []byte("Sneaky")), + hotline.NewField(hotline.FieldUserPassword, hotline.EncodeString([]byte("pw"))), + hotline.NewField(hotline.FieldUserAccess, make([]byte, 8)), + )) + assert.True(t, isError(reply), "guest lacks CreateUser access and must be rejected") + }) +} + +func TestE2E_DisconnectNotifiesPeers(t *testing.T) { + s := startE2EServer(t) + + alice := connectE2E(t, s.addr, "guest", "") + bob := connectE2E(t, s.addr, "guest", "") + + // Ensure both sessions are fully established (and thus both in the client manager) before alice + // leaves, so bob is guaranteed to be a peer that receives the delete notification. + require.False(t, isError(alice.roundTrip(newReq(hotline.TranGetUserNameList)))) + require.False(t, isError(bob.roundTrip(newReq(hotline.TranGetUserNameList)))) + + alice.cancel() + _ = alice.client.Disconnect() + + del := bob.waitFor(hotline.TranNotifyDeleteUser) + assert.Equal(t, hotline.TranNotifyDeleteUser, hotline.TranType(del.Type)) +} + +func TestE2E_ShutdownBroadcast(t *testing.T) { + if testing.Short() { + t.Skip("Shutdown pays a 3s broadcast-flush sleep") + } + s := startE2EServer(t) + c := connectE2E(t, s.addr, "guest", "") + // Ensure login has completed before triggering shutdown. + require.False(t, isError(c.roundTrip(newReq(hotline.TranGetUserNameList)))) + + go s.srv.Shutdown([]byte("server going down")) + + msg := c.waitFor(hotline.TranDisconnectMsg) + data := msg.GetField(hotline.FieldData) + require.NotNil(t, data) + assert.Contains(t, string(data.Data), "server going down") +} + +// --- upload payload construction ------------------------------------------------------------- + +func u32(v uint32) []byte { + b := make([]byte, 4) + binary.BigEndian.PutUint32(b, v) + return b +} + +// buildUploadFFO assembles the flattened file object bytes a client streams during an upload: +// FlatFileHeader + INFO fork header + info fork + DATA fork header + data fork. It mirrors the +// layout hotline.flattenedFileObject.ReadFrom expects (that type is unexported, so we build the +// bytes by hand from the exported information fork). +func buildUploadFFO(name string, content []byte) []byte { + infoFork := hotline.NewFlatFileInformationFork(name, [8]byte{}, "TEXT", "TTXT") + infoBody := mustReadAll(&infoFork) + + var buf bytes.Buffer + // FlatFileHeader: "FILP" + version 1 + 16 reserved + fork count 2. + buf.WriteString("FILP") + buf.Write([]byte{0, 1}) + buf.Write(make([]byte, 16)) + buf.Write([]byte{0, 2}) + + // INFO fork header + body. + buf.WriteString("INFO") + buf.Write(make([]byte, 8)) // compression + reserved + buf.Write(u32(uint32(len(infoBody)))) + buf.Write(infoBody) + + // DATA fork header + body. + buf.WriteString("DATA") + buf.Write(make([]byte, 8)) + buf.Write(u32(uint32(len(content)))) + buf.Write(content) + + return buf.Bytes() +} diff --git a/internal/mobius/integration_test.go b/internal/mobius/integration_test.go new file mode 100644 index 0000000..f1485eb --- /dev/null +++ b/internal/mobius/integration_test.go @@ -0,0 +1,429 @@ +package mobius + +// This file provides an in-process, protocol-level end-to-end harness: it builds a fully wired +// Hotline server (real YAML managers on a t.TempDir copy of test/config) on ephemeral ports and +// drives it with the in-repo hotline.Client. The individual regression tests live in e2e_test.go. + +import ( + "context" + "fmt" + "io" + "net" + "os" + "path" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/jhalter/mobius/hotline" + "github.com/stretchr/testify/require" + "golang.org/x/time/rate" +) + +// e2eServer is a running in-process server plus the metadata tests need to connect to it. +type e2eServer struct { + srv *hotline.Server + addr string // session port, host:port + xferAddr string // file transfer port (session port + 1), host:port + cfgDir string // tempdir holding the copied config + Files tree +} + +// startE2EServer builds and starts a fully-wired server on a free ephemeral port pair, returning +// once the session port accepts connections. The server is shut down via t.Cleanup. +// +// Binding is inherently racy: between probing for a free port pair and ListenAndServe claiming it, +// another process (a concurrently-tested package, or anything else on the machine) can steal a +// port. A stolen port surfaces as an early ListenAndServe error, so the serve attempt is retried +// on a fresh pair rather than failing the test. +func startE2EServer(t *testing.T) *e2eServer { + t.Helper() + + cfgDir := t.TempDir() + copyDirRecursiveTest(t, "test/config", cfgDir) + + cfg, err := LoadConfig(filepath.Join(cfgDir, "config.yaml")) + require.NoError(t, err) + // The fixture's FileRoot is a placeholder; point it at the copied Files tree. + cfg.FileRoot = filepath.Join(cfgDir, "Files") + + // Wire the concrete managers exactly as cmd/mobius-hotline-server/main.go does. + messageBoard, err := NewFlatNews(path.Join(cfgDir, "MessageBoard.txt")) + require.NoError(t, err) + + banFile, err := NewBanFile(path.Join(cfgDir, "Banlist.yaml")) + require.NoError(t, err) + + threadedNews, err := NewThreadedNewsYAML(path.Join(cfgDir, "ThreadedNews.yaml")) + require.NoError(t, err) + + am, err := NewYAMLAccountManager(path.Join(cfgDir, "Users/")) + require.NoError(t, err) + + agreement, err := NewAgreement(cfgDir, "\r") + require.NoError(t, err) + + // Create an admin account with a known password so tests that need elevated rights can log in + // deterministically. The client obfuscates passwords with EncodeString and the server bcrypt- + // compares that obfuscated form directly (see ClientConn.Authenticate), so the stored account + // must hash the obfuscated password — exactly what HandleNewUser does for real accounts. + var adminAccess hotline.AccessBitmap + for i := 0; i <= hotline.AccessSendPrivMsg; i++ { + adminAccess.Set(i) + } + obfuscatedPass := string(hotline.EncodeString([]byte(e2eAdminPass))) + require.NoError(t, am.Create(*hotline.NewAccount(e2eAdminLogin, "e2e admin", obfuscatedPass, adminAccess))) + + for attempt := 0; attempt < 5; attempt++ { + port := findFreePortPairTest(t) + + srv, err := hotline.NewServer( + hotline.WithInterface("127.0.0.1"), + hotline.WithPort(port), + hotline.WithConfig(*cfg), + hotline.WithLogger(NewTestLogger()), + // Disable per-IP connection throttling so multi-client tests don't pay 2s per connection. + hotline.WithConnectionRateLimit(rate.Inf, 1), + ) + require.NoError(t, err) + + srv.MessageBoard = messageBoard + srv.BanList = banFile + srv.ThreadedNewsMgr = threadedNews + srv.AccountManager = am + srv.Agreement = agreement + RegisterHandlers(srv) + + ctx, cancel := context.WithCancel(context.Background()) + serveErr := make(chan error, 1) + go func() { serveErr <- srv.ListenAndServe(ctx) }() + + addr := fmt.Sprintf("127.0.0.1:%d", port) + if !waitForListener(t, addr, serveErr) { + cancel() + continue + } + + t.Cleanup(func() { + cancel() + select { + case <-serveErr: + case <-time.After(10 * time.Second): + t.Error("server did not shut down within 10s") + } + }) + + return &e2eServer{ + srv: srv, + addr: addr, + xferAddr: fmt.Sprintf("127.0.0.1:%d", port+1), + cfgDir: cfgDir, + } + } + + t.Fatal("could not start e2e server: probed port pairs kept being claimed by other processes") + return nil +} + +const ( + e2eAdminLogin = "e2e-admin" + e2eAdminPass = "e2e-admin-pw" +) + +// waitForListener retry-dials until the server accepts a TCP connection, returning true on +// success. It returns false if ListenAndServe returned early instead — a bind failure, meaning +// another process claimed a probed port first — which the caller treats as a retry signal. A +// dial can even "succeed" against that other process's listener, so a successful dial only counts +// after the bind error has had a moment to surface. A server that neither accepts nor errors +// within the deadline fails the test. +func waitForListener(t *testing.T, addr string, serveErr <-chan error) bool { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + select { + case <-serveErr: + return false + default: + } + conn, err := net.DialTimeout("tcp", addr, 200*time.Millisecond) + if err == nil { + _ = conn.Close() + select { + case <-serveErr: + return false + case <-time.After(20 * time.Millisecond): + return true + } + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("server at %s never became ready", addr) + return false +} + +// findFreePortPairTest finds a port p such that both p and p+1 are free on the loopback interface. +// Duplicated from hotline/server_test.go's findFreePortPair (that copy is in-package and unexported). +func findFreePortPairTest(t *testing.T) int { + t.Helper() + for attempt := 0; attempt < 50; attempt++ { + l0, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + p := l0.Addr().(*net.TCPAddr).Port + l1, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", p+1)) + _ = l0.Close() + if err != nil { + continue + } + _ = l1.Close() + return p + } + t.Fatal("could not find a free consecutive port pair") + return 0 +} + +// copyDirRecursiveTest copies a directory tree from src into dst, creating dst subdirectories as +// needed. It mirrors copyDirRecursive in cmd/mobius-hotline-server (not importable from here). +func copyDirRecursiveTest(t *testing.T, src, dst string) { + t.Helper() + entries, err := os.ReadDir(src) + require.NoError(t, err) + require.NoError(t, os.MkdirAll(dst, 0755)) + + for _, entry := range entries { + srcPath := filepath.Join(src, entry.Name()) + dstPath := filepath.Join(dst, entry.Name()) + if entry.IsDir() { + copyDirRecursiveTest(t, srcPath, dstPath) + continue + } + data, err := os.ReadFile(srcPath) + require.NoError(t, err) + require.NoError(t, os.WriteFile(dstPath, data, 0644)) + } +} + +// --- E2E client driver ----------------------------------------------------------------------- + +// e2eClient wraps hotline.Client with request/reply correlation and collection of unsolicited +// server-pushed transactions, so tests can drive the protocol synchronously. +type e2eClient struct { + t *testing.T + client *hotline.Client + cancel context.CancelFunc + + mu sync.Mutex + waiters map[[4]byte]chan hotline.Transaction // reply ID -> waiter + incoming []hotline.Transaction // unsolicited (IsReply==0) transactions + incCh chan hotline.Transaction // signals a new unsolicited transaction +} + +// connectE2E performs handshake + login and starts the read loop. It auto-replies to the agreement +// prompt so the session becomes fully established. The login doubles as the display name. +func connectE2E(t *testing.T, addr, login, pass string) *e2eClient { + t.Helper() + return connectE2ENamed(t, addr, login, login, pass) +} + +// connectE2ENamed is connectE2E with a display name distinct from the login, so tests that run +// several sessions on the same account can tell them apart in user lists and notifications. +func connectE2ENamed(t *testing.T, addr, name, login, pass string) *e2eClient { + t.Helper() + + c := hotline.NewClient(name, NewTestLogger()) + ec := &e2eClient{ + t: t, + client: c, + waiters: map[[4]byte]chan hotline.Transaction{}, + incCh: make(chan hotline.Transaction, 64), + } + + // Register the same dispatcher for every transaction type the tests care about. Replies are + // re-typed to their request type by the client before dispatch, so one handler covers both + // request-reply and server-push flows. + for _, tt2 := range dispatchTypes { + c.HandleFunc(tt2, ec.dispatch) + } + + require.NoError(t, c.Connect(addr, login, pass)) + + ctx, cancel := context.WithCancel(context.Background()) + ec.cancel = cancel + go func() { _ = c.HandleTransactions(ctx) }() + + t.Cleanup(func() { + cancel() + _ = c.Disconnect() + }) + + return ec +} + +// dispatchTypes is the set of transaction types the e2e dispatcher is registered for: every request +// type the tests send (replies arrive re-typed to these) plus server-pushed types. +var dispatchTypes = [][2]byte{ + hotline.TranLogin, hotline.TranAgreed, hotline.TranShowAgreement, + hotline.TranChatSend, hotline.TranChatMsg, + hotline.TranGetMsgs, hotline.TranOldPostNews, + hotline.TranGetNewsCatNameList, hotline.TranGetNewsArtNameList, + hotline.TranGetNewsArtData, hotline.TranPostNewsArt, + hotline.TranGetFileNameList, hotline.TranDownloadFile, hotline.TranUploadFile, + hotline.TranGetUser, hotline.TranNewUser, hotline.TranSetUser, hotline.TranDeleteUser, + hotline.TranListUsers, hotline.TranGetUserNameList, + hotline.TranNotifyChangeUser, hotline.TranNotifyDeleteUser, + hotline.TranServerMsg, hotline.TranDisconnectMsg, hotline.TranUserAccess, + hotline.TranInviteNewChat, hotline.TranInviteToChat, hotline.TranJoinChat, + hotline.TranNotifyChatChangeUser, hotline.TranNotifyChatDeleteUser, hotline.TranNotifyChatSubject, +} + +func (ec *e2eClient) dispatch(_ context.Context, _ *hotline.Client, t *hotline.Transaction) ([]hotline.Transaction, error) { + if t.IsReply == 1 { + ec.mu.Lock() + ch := ec.waiters[t.ID] + delete(ec.waiters, t.ID) + ec.mu.Unlock() + if ch != nil { + ch <- *t + } + return nil, nil + } + + // Unsolicited server push. + if t.Type == hotline.TranShowAgreement { + // Accept the agreement so login completes. + return []hotline.Transaction{ + hotline.NewTransaction(hotline.TranAgreed, [2]byte{}, + hotline.NewField(hotline.FieldUserName, []byte(ec.client.Pref.Username)), + hotline.NewField(hotline.FieldUserIconID, ec.client.Pref.IconBytes()), + hotline.NewField(hotline.FieldOptions, []byte{0, 0}), + ), + }, nil + } + + ec.mu.Lock() + ec.incoming = append(ec.incoming, *t) + ec.mu.Unlock() + select { + case ec.incCh <- *t: + default: + } + return nil, nil +} + +// roundTrip sends a request and waits for the matching reply (correlated by transaction ID). +func (ec *e2eClient) roundTrip(req hotline.Transaction) hotline.Transaction { + ec.t.Helper() + + ch := make(chan hotline.Transaction, 1) + ec.mu.Lock() + ec.waiters[req.ID] = ch + ec.mu.Unlock() + + require.NoError(ec.t, ec.client.Send(req)) + + select { + case reply := <-ch: + return reply + case <-time.After(15 * time.Second): + ec.t.Fatalf("timed out waiting for reply to %v", req.Type) + return hotline.Transaction{} + } +} + +// waitFor blocks until an unsolicited transaction of the given type arrives, or the test fails. +// It polls the buffer on a short interval as well as waking on incCh, so a signal that races with +// the buffer scan can never cause a lost wakeup. +func (ec *e2eClient) waitFor(tranType [2]byte) hotline.Transaction { + ec.t.Helper() + + deadline := time.After(15 * time.Second) + ticker := time.NewTicker(20 * time.Millisecond) + defer ticker.Stop() + + for { + // Scan anything already buffered. + ec.mu.Lock() + for i, tr := range ec.incoming { + if tr.Type == tranType { + ec.incoming = append(ec.incoming[:i], ec.incoming[i+1:]...) + ec.mu.Unlock() + return tr + } + } + ec.mu.Unlock() + + select { + case <-ec.incCh: + // Woke on a new arrival; loop and re-scan. + case <-ticker.C: + // Periodic re-scan guards against a dropped incCh signal. + case <-deadline: + ec.t.Fatalf("timed out waiting for unsolicited %v", tranType) + return hotline.Transaction{} + } + } +} + +// isError reports whether a reply carries a Hotline error (non-zero ErrorCode). +func isError(t hotline.Transaction) bool { + return t.ErrorCode != [4]byte{0, 0, 0, 0} +} + +// --- transfer-port helpers ------------------------------------------------------------------- + +// htxfHeader builds the 16-byte HTXF transfer header (protocol + reference number + data size). +func htxfHeader(refNum [4]byte, dataSize uint32) []byte { + b := make([]byte, 16) + copy(b[0:4], hotline.HTXF[:]) + copy(b[4:8], refNum[:]) + b[8] = byte(dataSize >> 24) + b[9] = byte(dataSize >> 16) + b[10] = byte(dataSize >> 8) + b[11] = byte(dataSize) + return b +} + +// downloadOverTransferPort dials the transfer port, sends the HTXF header for refNum, and returns +// all bytes the server streams back (the flattened file object). +func downloadOverTransferPort(t *testing.T, xferAddr string, refNum [4]byte) []byte { + t.Helper() + conn, err := net.DialTimeout("tcp", xferAddr, 5*time.Second) + require.NoError(t, err) + defer func() { _ = conn.Close() }() + + _, err = conn.Write(htxfHeader(refNum, 0)) + require.NoError(t, err) + + require.NoError(t, conn.SetReadDeadline(time.Now().Add(5*time.Second))) + data, err := io.ReadAll(conn) + // A read deadline or server-initiated close both surface here; the payload collected so far is + // what we assert on. + if err != nil && !isTimeout(err) { + require.ErrorIs(t, err, io.EOF) + } + return data +} + +// uploadOverTransferPort dials the transfer port, sends the HTXF header, then streams payload. +func uploadOverTransferPort(t *testing.T, xferAddr string, refNum [4]byte, payload []byte) { + t.Helper() + conn, err := net.DialTimeout("tcp", xferAddr, 5*time.Second) + require.NoError(t, err) + defer func() { _ = conn.Close() }() + + _, err = conn.Write(htxfHeader(refNum, uint32(len(payload)))) + require.NoError(t, err) + _, err = conn.Write(payload) + require.NoError(t, err) + + // Give the server a moment to persist before the caller asserts. + require.NoError(t, conn.SetReadDeadline(time.Now().Add(2*time.Second))) + _, _ = io.ReadAll(conn) +} + +func isTimeout(err error) bool { + var ne net.Error + if e, ok := err.(net.Error); ok { + ne = e + } + return ne != nil && ne.Timeout() +} diff --git a/internal/mobius/reload_test.go b/internal/mobius/reload_test.go new file mode 100644 index 0000000..d864fa7 --- /dev/null +++ b/internal/mobius/reload_test.go @@ -0,0 +1,32 @@ +package mobius + +import ( + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// The concrete Reload implementations (FlatNews, Agreement, BanFile, ThreadedNewsYAML) are +// exercised in their own *_test.go files. This covers the ReloaderFunc adapter, which is the +// only unit defined in reload.go without direct coverage. +func TestReloaderFunc_Reload(t *testing.T) { + t.Run("invokes the wrapped func", func(t *testing.T) { + called := false + var r Reloader = ReloaderFunc(func() error { + called = true + return nil + }) + + require.NoError(t, r.Reload()) + assert.True(t, called) + }) + + t.Run("propagates the wrapped error", func(t *testing.T) { + sentinel := errors.New("reload failed") + var r Reloader = ReloaderFunc(func() error { return sentinel }) + + assert.ErrorIs(t, r.Reload(), sentinel) + }) +} |