From 2f46f87177234070044b5304ca3e0db471699ad8 Mon Sep 17 00:00:00 2001 From: Jeff Halter <868228+jhalter@users.noreply.github.com> Date: Mon, 16 Mar 2026 11:46:04 -0700 Subject: Improve test coverage for hotline and internal/mobius packages MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add comprehensive test cases across both packages to increase coverage: - hotline: 52.4% → ~55% (Disconnect, handleTransaction, SendAll, sendBanMessage, MemClientMgr, and other tests) - internal/mobius: 75.9% → ~80% (HandleUpdateUser, HandleDeleteUser, HandleSetUser, HandleNewUser, HandleUserBroadcast success paths) --- hotline/account_test.go | 91 ++++++ hotline/ban.go | 70 ++++- hotline/client_conn_test.go | 505 ++++++++++++++++++++++++++++++++++ hotline/client_manager_test.go | 62 +++++ hotline/client_test.go | 278 +++++++++++++++++++ hotline/file_wrapper_test.go | 171 ++++++++++++ hotline/flattened_file_object_test.go | 142 ++++++++++ hotline/server_test.go | 110 ++++++++ 8 files changed, 1428 insertions(+), 1 deletion(-) create mode 100644 hotline/account_test.go create mode 100644 hotline/client_manager_test.go create mode 100644 hotline/client_test.go create mode 100644 hotline/file_wrapper_test.go (limited to 'hotline') diff --git a/hotline/account_test.go b/hotline/account_test.go new file mode 100644 index 0000000..43fe28f --- /dev/null +++ b/hotline/account_test.go @@ -0,0 +1,91 @@ +package hotline + +import ( + "encoding/binary" + "io" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/bcrypt" +) + +func TestNewAccount(t *testing.T) { + access := AccessBitmap{0xff, 0, 0, 0, 0, 0, 0, 0} + acct := NewAccount("jdoe", "John Doe", "secret123", access) + + assert.Equal(t, "jdoe", acct.Login) + assert.Equal(t, "John Doe", acct.Name) + assert.Equal(t, access, acct.Access) + + // Password should be a bcrypt hash, not the plaintext. + assert.NotEqual(t, "secret123", acct.Password) + err := bcrypt.CompareHashAndPassword([]byte(acct.Password), []byte("secret123")) + require.NoError(t, err, "password hash should validate against original plaintext") +} + +func TestHashAndSalt(t *testing.T) { + t.Run("produces valid bcrypt hash", func(t *testing.T) { + hash := HashAndSalt([]byte("password")) + err := bcrypt.CompareHashAndPassword([]byte(hash), []byte("password")) + require.NoError(t, err) + }) + + t.Run("different calls produce different hashes", func(t *testing.T) { + h1 := HashAndSalt([]byte("same")) + h2 := HashAndSalt([]byte("same")) + assert.NotEqual(t, h1, h2, "bcrypt salts should differ between calls") + + // Both should still verify against the original input. + require.NoError(t, bcrypt.CompareHashAndPassword([]byte(h1), []byte("same"))) + require.NoError(t, bcrypt.CompareHashAndPassword([]byte(h2), []byte("same"))) + }) +} + +func TestAccount_Read(t *testing.T) { + t.Run("with password set", func(t *testing.T) { + acct := NewAccount("admin", "Admin User", "pass", AccessBitmap{}) + + data, err := io.ReadAll(acct) + require.NoError(t, err) + + // First two bytes are the field count (big-endian uint16). + require.GreaterOrEqual(t, len(data), 2) + fieldCount := binary.BigEndian.Uint16(data[:2]) + assert.Equal(t, uint16(4), fieldCount, "should have 4 fields when password is set") + + // Verify the password marker "x" appears in the serialized data. + assert.Contains(t, string(data), "x") + + // Verify the user name appears in the serialized data. + assert.Contains(t, string(data), "Admin User") + }) + + t.Run("with empty password", func(t *testing.T) { + acct := &Account{ + Login: "guest", + Name: "Guest", + Password: HashAndSalt([]byte("")), + Access: AccessBitmap{}, + } + + data, err := io.ReadAll(acct) + require.NoError(t, err) + + require.GreaterOrEqual(t, len(data), 2) + fieldCount := binary.BigEndian.Uint16(data[:2]) + assert.Equal(t, uint16(3), fieldCount, "should have 3 fields when password is empty") + }) + + t.Run("full read returns all serialized bytes", func(t *testing.T) { + acct := NewAccount("test", "Test", "pw", AccessBitmap{0x01}) + + data, err := io.ReadAll(acct) + require.NoError(t, err) + assert.Greater(t, len(data), 2, "serialized output should contain field data beyond the count header") + }) +} + +func TestNewAccount_GuestConstant(t *testing.T) { + assert.Equal(t, "guest", GuestAccount) +} diff --git a/hotline/ban.go b/hotline/ban.go index cc2ecde..2713fbe 100644 --- a/hotline/ban.go +++ b/hotline/ban.go @@ -1,6 +1,10 @@ package hotline -import "time" +import ( + "time" + + "github.com/stretchr/testify/mock" +) // BanDuration is the length of time for temporary bans. const BanDuration = 30 * time.Minute @@ -36,3 +40,67 @@ type BanMgr interface { ListBannedUsernames() ([]string, error) ListBannedNicknames() ([]string, error) } + +type MockBanMgr struct { + mock.Mock +} + +func (m *MockBanMgr) Add(ip string, until *time.Time) error { + args := m.Called(ip, until) + return args.Error(0) +} + +func (m *MockBanMgr) IsBanned(ip string) (bool, *time.Time) { + args := m.Called(ip) + return args.Bool(0), args.Get(1).(*time.Time) +} + +func (m *MockBanMgr) UnbanIP(ip string) error { + args := m.Called(ip) + return args.Error(0) +} + +func (m *MockBanMgr) BanUsername(username string) error { + args := m.Called(username) + return args.Error(0) +} + +func (m *MockBanMgr) UnbanUsername(username string) error { + args := m.Called(username) + return args.Error(0) +} + +func (m *MockBanMgr) IsUsernameBanned(username string) bool { + args := m.Called(username) + return args.Bool(0) +} + +func (m *MockBanMgr) BanNickname(nickname string) error { + args := m.Called(nickname) + return args.Error(0) +} + +func (m *MockBanMgr) UnbanNickname(nickname string) error { + args := m.Called(nickname) + return args.Error(0) +} + +func (m *MockBanMgr) IsNicknameBanned(nickname string) bool { + args := m.Called(nickname) + return args.Bool(0) +} + +func (m *MockBanMgr) ListBannedIPs() ([]string, error) { + args := m.Called() + return args.Get(0).([]string), args.Error(1) +} + +func (m *MockBanMgr) ListBannedUsernames() ([]string, error) { + args := m.Called() + return args.Get(0).([]string), args.Error(1) +} + +func (m *MockBanMgr) ListBannedNicknames() ([]string, error) { + args := m.Called() + return args.Get(0).([]string), args.Error(1) +} diff --git a/hotline/client_conn_test.go b/hotline/client_conn_test.go index e5acf82..94e38e9 100644 --- a/hotline/client_conn_test.go +++ b/hotline/client_conn_test.go @@ -1 +1,506 @@ package hotline + +import ( + "bytes" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/bcrypt" +) + +type mockAccountMgr struct { + accounts map[string]*Account +} + +func (m *mockAccountMgr) Create(account Account) error { return nil } +func (m *mockAccountMgr) Update(account Account, newLogin string) error { return nil } +func (m *mockAccountMgr) Get(login string) *Account { return m.accounts[login] } +func (m *mockAccountMgr) List() []Account { return nil } +func (m *mockAccountMgr) Delete(login string) error { return nil } + +func TestClientConn_IP(t *testing.T) { + tests := []struct { + name string + remoteAddr string + wantIP string + }{ + { + name: "extracts IP from host:port", + remoteAddr: "192.168.1.1:12345", + wantIP: "192.168.1.1", + }, + { + name: "extracts IPv6 address", + remoteAddr: "[::1]:12345", + wantIP: "::1", + }, + { + name: "returns empty string for missing port", + remoteAddr: "192.168.1.1", + wantIP: "", + }, + { + name: "returns empty string for empty RemoteAddr", + remoteAddr: "", + wantIP: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cc := &ClientConn{RemoteAddr: tt.remoteAddr} + assert.Equal(t, tt.wantIP, cc.IP()) + }) + } +} + +func TestClientConn_FileRoot(t *testing.T) { + tests := []struct { + name string + accountFileRoot string + serverFileRoot string + want string + }{ + { + name: "returns account FileRoot when set", + accountFileRoot: "/home/user/files", + serverFileRoot: "/srv/files", + want: "/home/user/files", + }, + { + name: "falls back to server Config FileRoot", + accountFileRoot: "", + serverFileRoot: "/srv/files", + want: "/srv/files", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cc := &ClientConn{ + Account: &Account{FileRoot: tt.accountFileRoot}, + Server: &Server{Config: Config{FileRoot: tt.serverFileRoot}}, + } + assert.Equal(t, tt.want, cc.FileRoot()) + }) + } +} + +func TestClientConn_Authenticate(t *testing.T) { + password := []byte("secret123") + hash, err := bcrypt.GenerateFromPassword(password, bcrypt.MinCost) + require.NoError(t, err) + + mgr := &mockAccountMgr{ + accounts: map[string]*Account{ + "admin": {Login: "admin", Password: string(hash)}, + }, + } + + tests := []struct { + name string + login string + password []byte + want bool + }{ + { + name: "valid credentials return true", + login: "admin", + password: password, + want: true, + }, + { + name: "wrong password returns false", + login: "admin", + password: []byte("wrongpassword"), + want: false, + }, + { + name: "nonexistent login returns false", + login: "nobody", + password: password, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cc := &ClientConn{ + Server: &Server{AccountManager: mgr}, + } + assert.Equal(t, tt.want, cc.Authenticate(tt.login, tt.password)) + }) + } +} + +func TestClientConn_Authorize(t *testing.T) { + tests := []struct { + name string + account *Account + access int + want bool + }{ + { + name: "returns false when Account is nil", + account: nil, + access: AccessDeleteFile, + want: false, + }, + { + name: "returns true when access bit is set", + account: func() *Account { + a := &Account{} + a.Access.Set(AccessUploadFile) + return a + }(), + access: AccessUploadFile, + want: true, + }, + { + name: "returns false when access bit is not set", + account: &Account{}, + access: AccessDownloadFile, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cc := &ClientConn{Account: tt.account} + assert.Equal(t, tt.want, cc.Authorize(tt.access)) + }) + } +} + +func TestClientConn_NewReply(t *testing.T) { + tests := []struct { + name string + clientID ClientID + tranID [4]byte + fields []Field + }{ + { + name: "reply with no fields", + clientID: ClientID{0x00, 0x01}, + tranID: [4]byte{0x00, 0x00, 0x00, 0x05}, + fields: nil, + }, + { + name: "reply with fields", + clientID: ClientID{0x00, 0x02}, + tranID: [4]byte{0x00, 0x00, 0x00, 0x0A}, + fields: []Field{NewField(FieldError, []byte("test"))}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cc := &ClientConn{ID: tt.clientID} + tran := &Transaction{ID: tt.tranID} + + reply := cc.NewReply(tran, tt.fields...) + + assert.Equal(t, byte(1), reply.IsReply) + assert.Equal(t, tt.tranID, reply.ID) + assert.Equal(t, tt.clientID, reply.ClientID) + assert.Equal(t, tt.fields, reply.Fields) + }) + } +} + +func TestClientConn_NewErrReply(t *testing.T) { + tests := []struct { + name string + clientID ClientID + tranID [4]byte + errMsg string + }{ + { + name: "returns error transaction", + clientID: ClientID{0x00, 0x01}, + tranID: [4]byte{0x00, 0x00, 0x00, 0x01}, + errMsg: "access denied", + }, + { + name: "returns error with empty message", + clientID: ClientID{0x00, 0x03}, + tranID: [4]byte{0x00, 0x00, 0x00, 0x02}, + errMsg: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cc := &ClientConn{ID: tt.clientID} + tran := &Transaction{ID: tt.tranID} + + result := cc.NewErrReply(tran, tt.errMsg) + + require.Len(t, result, 1) + errReply := result[0] + + assert.Equal(t, byte(1), errReply.IsReply) + assert.Equal(t, tt.tranID, errReply.ID) + assert.Equal(t, tt.clientID, errReply.ClientID) + assert.Equal(t, [4]byte{0, 0, 0, 1}, errReply.ErrorCode) + + require.Len(t, errReply.Fields, 1) + assert.Equal(t, FieldError, errReply.Fields[0].Type) + assert.Equal(t, []byte(tt.errMsg), errReply.Fields[0].Data) + }) + } +} + +func TestClientConn_NotifyOthers(t *testing.T) { + tests := []struct { + name string + selfID ClientID + otherConns []*ClientConn + wantCount int + wantClients []ClientID + }{ + { + name: "sends to other clients, excludes self", + selfID: ClientID{0x00, 0x01}, + otherConns: []*ClientConn{ + {ID: ClientID{0x00, 0x01}}, + {ID: ClientID{0x00, 0x02}}, + {ID: ClientID{0x00, 0x03}}, + }, + wantCount: 2, + wantClients: []ClientID{{0x00, 0x02}, {0x00, 0x03}}, + }, + { + name: "returns nil when only self is connected", + selfID: ClientID{0x00, 0x01}, + otherConns: []*ClientConn{ + {ID: ClientID{0x00, 0x01}}, + }, + wantCount: 0, + wantClients: nil, + }, + { + name: "returns nil when no clients are connected", + selfID: ClientID{0x00, 0x01}, + otherConns: []*ClientConn{}, + wantCount: 0, + wantClients: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + mockMgr := &MockClientMgr{} + mockMgr.On("List").Return(tt.otherConns) + + cc := &ClientConn{ + ID: tt.selfID, + Server: &Server{ClientMgr: mockMgr}, + } + + tran := NewTransaction(TranChatSend, ClientID{}, NewField(FieldError, []byte("test"))) + result := cc.NotifyOthers(tran) + + assert.Len(t, result, tt.wantCount) + + if tt.wantClients != nil { + var gotIDs []ClientID + for _, r := range result { + gotIDs = append(gotIDs, r.ClientID) + } + assert.Equal(t, tt.wantClients, gotIDs) + } + + mockMgr.AssertExpectations(t) + }) + } +} + +func TestClientConn_String(t *testing.T) { + cc := &ClientConn{ + UserName: []byte("TestUser"), + RemoteAddr: "192.168.1.10:54321", + Account: &Account{Name: "Test Account", Login: "testlogin"}, + ClientFileTransferMgr: NewClientFileTransferMgr(), + } + + result := cc.String() + + // The method replaces \n with \r + assert.Contains(t, result, "TestUser") + assert.Contains(t, result, "Test Account") + assert.Contains(t, result, "testlogin") + assert.Contains(t, result, "192.168.1.10:54321") + assert.Contains(t, result, "None.") + assert.NotContains(t, result, "\n", "all newlines should be replaced with carriage returns") + assert.Contains(t, result, "\r") +} + +func TestClientConn_Disconnect(t *testing.T) { + t.Run("removes client and closes connection", func(t *testing.T) { + mockMgr := &MockClientMgr{} + mockMgr.On("Delete", ClientID{0, 1}).Return() + mockMgr.On("List").Return([]*ClientConn{}) + + outbox := make(chan Transaction, 10) + cc := &ClientConn{ + ID: ClientID{0, 1}, + Connection: &nopCloserRWC{Buffer: &bytes.Buffer{}}, + Server: &Server{ + ClientMgr: mockMgr, + outbox: outbox, + Logger: NewTestLogger(), + }, + } + + cc.Disconnect() + + mockMgr.AssertCalled(t, "Delete", ClientID{0, 1}) + assert.Empty(t, outbox) // No other clients to notify + }) + + t.Run("notifies other clients", func(t *testing.T) { + mockMgr := &MockClientMgr{} + mockMgr.On("Delete", ClientID{0, 1}).Return() + mockMgr.On("List").Return([]*ClientConn{ + {ID: ClientID{0, 1}}, + {ID: ClientID{0, 2}}, + {ID: ClientID{0, 3}}, + }) + + outbox := make(chan Transaction, 10) + cc := &ClientConn{ + ID: ClientID{0, 1}, + Connection: &nopCloserRWC{Buffer: &bytes.Buffer{}}, + Server: &Server{ + ClientMgr: mockMgr, + outbox: outbox, + Logger: NewTestLogger(), + }, + } + + cc.Disconnect() + + assert.Len(t, outbox, 2) + mockMgr.AssertExpectations(t) + }) +} + +func TestClientConn_handleTransaction(t *testing.T) { + t.Run("dispatches to registered handler", func(t *testing.T) { + outbox := make(chan Transaction, 10) + mockMgr := &MockClientMgr{} + + cc := &ClientConn{ + ID: ClientID{0, 1}, + Account: &Account{}, + Logger: NewTestLogger(), + Server: &Server{ + outbox: outbox, + ClientMgr: mockMgr, + handlers: map[TranType]HandlerFunc{ + TranChatSend: func(cc *ClientConn, t *Transaction) []Transaction { + return []Transaction{cc.NewReply(t)} + }, + }, + }, + } + + cc.handleTransaction(NewTransaction(TranChatSend, ClientID{0, 1})) + + assert.Len(t, outbox, 1) + assert.Equal(t, 0, cc.IdleTime) + }) + + t.Run("keepalive does not reset idle time", func(t *testing.T) { + outbox := make(chan Transaction, 10) + + cc := &ClientConn{ + ID: ClientID{0, 1}, + Account: &Account{}, + IdleTime: 100, + Logger: NewTestLogger(), + Server: &Server{ + outbox: outbox, + handlers: map[TranType]HandlerFunc{ + TranKeepAlive: func(cc *ClientConn, t *Transaction) []Transaction { + return []Transaction{cc.NewReply(t)} + }, + }, + }, + } + + cc.handleTransaction(NewTransaction(TranKeepAlive, ClientID{0, 1})) + + assert.Equal(t, 100, cc.IdleTime) + }) + + t.Run("non-keepalive clears away flag", func(t *testing.T) { + outbox := make(chan Transaction, 10) + mockMgr := &MockClientMgr{} + mockMgr.On("List").Return([]*ClientConn{ + {ID: ClientID{0, 1}}, + }) + + cc := &ClientConn{ + ID: ClientID{0, 1}, + Account: &Account{}, + UserName: []byte("test"), + Icon: []byte{0, 0}, + IdleTime: 50, + Logger: NewTestLogger(), + Server: &Server{ + outbox: outbox, + ClientMgr: mockMgr, + handlers: map[TranType]HandlerFunc{ + TranChatSend: func(cc *ClientConn, t *Transaction) []Transaction { + return nil + }, + }, + }, + } + cc.Flags.Set(UserFlagAway, 1) + + cc.handleTransaction(NewTransaction(TranChatSend, ClientID{0, 1})) + + assert.Equal(t, 0, cc.IdleTime) + assert.False(t, cc.Flags.IsSet(UserFlagAway)) + // SendAll should have sent TranNotifyChangeUser + assert.Greater(t, len(outbox), 0) + }) +} + +func TestClientConn_SendAll(t *testing.T) { + mockMgr := &MockClientMgr{} + mockMgr.On("List").Return([]*ClientConn{ + {ID: ClientID{0, 1}}, + {ID: ClientID{0, 2}}, + {ID: ClientID{0, 3}}, + }) + + outbox := make(chan Transaction, 10) + cc := &ClientConn{ + ID: ClientID{0, 1}, + Server: &Server{ + ClientMgr: mockMgr, + outbox: outbox, + }, + } + + cc.SendAll(TranChatMsg, NewField(FieldData, []byte("hello"))) + + assert.Len(t, outbox, 3) + + clientIDs := make(map[ClientID]bool) + for range 3 { + tran := <-outbox + clientIDs[tran.ClientID] = true + assert.Equal(t, TranChatMsg, tran.Type) + } + assert.True(t, clientIDs[ClientID{0, 1}]) + assert.True(t, clientIDs[ClientID{0, 2}]) + assert.True(t, clientIDs[ClientID{0, 3}]) + + mockMgr.AssertExpectations(t) +} diff --git a/hotline/client_manager_test.go b/hotline/client_manager_test.go new file mode 100644 index 0000000..a2a0320 --- /dev/null +++ b/hotline/client_manager_test.go @@ -0,0 +1,62 @@ +package hotline + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestMemClientMgr_Add(t *testing.T) { + mgr := NewMemClientMgr() + + c1 := &ClientConn{} + c2 := &ClientConn{} + mgr.Add(c1) + mgr.Add(c2) + + assert.NotEqual(t, c1.ID, c2.ID) + assert.Len(t, mgr.List(), 2) +} + +func TestMemClientMgr_Get(t *testing.T) { + mgr := NewMemClientMgr() + + assert.Nil(t, mgr.Get(ClientID{0xFF, 0xFF})) + + c := &ClientConn{} + mgr.Add(c) + assert.Equal(t, c, mgr.Get(c.ID)) +} + +func TestMemClientMgr_Delete(t *testing.T) { + mgr := NewMemClientMgr() + + c := &ClientConn{} + mgr.Add(c) + id := c.ID + + mgr.Delete(id) + assert.Nil(t, mgr.Get(id)) +} + +func TestMemClientMgr_List_Sorted(t *testing.T) { + mgr := NewMemClientMgr() + + // Add 3 clients - they'll get sequential IDs + c1 := &ClientConn{} + c2 := &ClientConn{} + c3 := &ClientConn{} + mgr.Add(c1) + mgr.Add(c2) + mgr.Add(c3) + + list := mgr.List() + assert.Len(t, list, 3) + + // Verify sorted by ID + for i := 1; i < len(list); i++ { + assert.True(t, list[i-1].ID[0] < list[i].ID[0] || + (list[i-1].ID[0] == list[i].ID[0] && list[i-1].ID[1] <= list[i].ID[1]), + "clients should be sorted by ID") + } +} diff --git a/hotline/client_test.go b/hotline/client_test.go new file mode 100644 index 0000000..23b6155 --- /dev/null +++ b/hotline/client_test.go @@ -0,0 +1,278 @@ +package hotline + +import ( + "bytes" + "context" + "io" + "log/slog" + "net" + "os" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func newTestClient() *Client { + return NewClient("testuser", slog.New(slog.NewTextHandler(os.Stdout, nil))) +} + +func TestClient_Handshake(t *testing.T) { + t.Run("successful handshake", func(t *testing.T) { + clientConn, serverConn := net.Pipe() + defer clientConn.Close() + defer serverConn.Close() + + c := newTestClient() + c.Connection = clientConn + + // Server side: read the client handshake and respond + go func() { + buf := make([]byte, 12) + _, _ = io.ReadFull(serverConn, buf) + assert.Equal(t, ClientHandshake, buf) + _, _ = serverConn.Write(ServerHandshake) + }() + + err := c.Handshake() + assert.NoError(t, err) + }) + + t.Run("server returns error response", func(t *testing.T) { + clientConn, serverConn := net.Pipe() + defer clientConn.Close() + defer serverConn.Close() + + c := newTestClient() + c.Connection = clientConn + + go func() { + buf := make([]byte, 12) + _, _ = io.ReadFull(serverConn, buf) + // Send a non-standard response + _, _ = serverConn.Write([]byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01}) + }() + + err := c.Handshake() + assert.Error(t, err) + assert.Contains(t, err.Error(), "unexpected handshake response") + }) + + t.Run("connection closed during read", func(t *testing.T) { + clientConn, serverConn := net.Pipe() + defer clientConn.Close() + + c := newTestClient() + c.Connection = clientConn + + go func() { + buf := make([]byte, 12) + _, _ = io.ReadFull(serverConn, buf) + serverConn.Close() + }() + + err := c.Handshake() + assert.Error(t, err) + assert.Contains(t, err.Error(), "handshake read err") + }) +} + +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() + + c := newTestClient() + c.Connection = clientConn + + tran := NewTransaction(TranChatSend, [2]byte{0, 1}, NewField(FieldData, []byte("hello"))) + + // Read what client sends + go func() { + buf := make([]byte, 4096) + _, _ = serverConn.Read(buf) + }() + + err := c.Send(tran) + assert.NoError(t, err) + + // Verify the transaction was added to activeTasks + assert.Contains(t, c.activeTasks, tran.ID) + }) + + t.Run("reply transactions are not tracked", func(t *testing.T) { + clientConn, serverConn := net.Pipe() + defer clientConn.Close() + defer serverConn.Close() + + c := newTestClient() + c.Connection = clientConn + + tran := NewTransaction(TranChatSend, [2]byte{0, 1}) + tran.IsReply = 1 + + go func() { + buf := make([]byte, 4096) + _, _ = serverConn.Read(buf) + }() + + err := c.Send(tran) + assert.NoError(t, err) + assert.NotContains(t, c.activeTasks, tran.ID) + }) +} + +func TestClient_HandleTransaction(t *testing.T) { + t.Run("dispatches to registered handler", func(t *testing.T) { + c := newTestClient() + c.Connection = &clientMockConn{WBuf: &bytes.Buffer{}} + + handlerCalled := false + c.HandleFunc(TranChatMsg, func(ctx context.Context, client *Client, t *Transaction) ([]Transaction, error) { + handlerCalled = true + return nil, nil + }) + + tran := &Transaction{Type: TranChatMsg} + err := c.HandleTransaction(context.Background(), tran) + + assert.NoError(t, err) + assert.True(t, handlerCalled) + }) + + t.Run("reply matches original request type", func(t *testing.T) { + c := newTestClient() + c.Connection = &clientMockConn{WBuf: &bytes.Buffer{}} + + // Register a handler for the original type + var receivedType TranType + c.HandleFunc(TranGetUserNameList, func(ctx context.Context, client *Client, t *Transaction) ([]Transaction, error) { + receivedType = t.Type + return nil, nil + }) + + // Simulate sending a request first + origTran := NewTransaction(TranGetUserNameList, [2]byte{0, 1}) + c.activeTasks[origTran.ID] = &origTran + + // Create a reply with the same ID + reply := &Transaction{ + ID: origTran.ID, + IsReply: 1, + } + + err := c.HandleTransaction(context.Background(), reply) + assert.NoError(t, err) + assert.Equal(t, TranGetUserNameList, receivedType) + }) + + t.Run("reply with no matching request returns error", func(t *testing.T) { + c := newTestClient() + c.Connection = &clientMockConn{WBuf: &bytes.Buffer{}} + + reply := &Transaction{ + ID: [4]byte{0, 0, 0, 99}, + IsReply: 1, + } + + err := c.HandleTransaction(context.Background(), reply) + assert.Error(t, err) + assert.Contains(t, err.Error(), "no matching request") + }) + + t.Run("unhandled transaction type does not error", func(t *testing.T) { + c := newTestClient() + c.Connection = &clientMockConn{WBuf: &bytes.Buffer{}} + + tran := &Transaction{Type: [2]byte{0xFF, 0xFF}} + err := c.HandleTransaction(context.Background(), tran) + assert.NoError(t, err) + }) +} + +func TestClient_Disconnect(t *testing.T) { + t.Run("closes connection and done channel", func(t *testing.T) { + clientConn, serverConn := net.Pipe() + defer serverConn.Close() + + c := newTestClient() + c.Connection = clientConn + c.done = make(chan struct{}) + doneCh := c.done // save reference before Disconnect sets it to nil + + err := c.Disconnect() + assert.NoError(t, err) + + // Verify done channel is closed + select { + case <-doneCh: + // expected - channel is closed + default: + t.Error("done channel should be closed") + } + + // Verify connection is closed by attempting a write + _, err = clientConn.Write([]byte("test")) + assert.Error(t, err) + }) + + t.Run("handles nil done channel", func(t *testing.T) { + clientConn, serverConn := net.Pipe() + defer serverConn.Close() + + c := newTestClient() + c.Connection = clientConn + // done is nil by default (before Connect is called) + + err := c.Disconnect() + assert.NoError(t, err) + }) +} + +func TestClient_HandleFunc(t *testing.T) { + c := newTestClient() + + handler := func(ctx context.Context, client *Client, t *Transaction) ([]Transaction, error) { + return nil, nil + } + + c.HandleFunc(TranChatMsg, handler) + + _, ok := c.Handlers[TranChatMsg] + require.True(t, ok, "handler should be registered") +} + +func TestNewClient(t *testing.T) { + logger := slog.New(slog.NewTextHandler(os.Stdout, nil)) + c := NewClient("myuser", logger) + + assert.Equal(t, "myuser", c.Pref.Username) + assert.NotNil(t, c.Handlers) + assert.NotNil(t, c.activeTasks) +} + +// clientMockConn implements net.Conn for testing Send without a net.Pipe. +type clientMockConn struct { + RBuf *bytes.Buffer + WBuf *bytes.Buffer +} + +func (mc *clientMockConn) Read(b []byte) (n int, err error) { + if mc.RBuf == nil { + return 0, io.EOF + } + return mc.RBuf.Read(b) +} + +func (mc *clientMockConn) Write(b []byte) (n int, err error) { + return mc.WBuf.Write(b) +} + +func (mc *clientMockConn) Close() error { return nil } +func (mc *clientMockConn) LocalAddr() net.Addr { return nil } +func (mc *clientMockConn) RemoteAddr() net.Addr { return nil } +func (mc *clientMockConn) SetDeadline(_ time.Time) error { return nil } +func (mc *clientMockConn) SetReadDeadline(_ time.Time) error { return nil } +func (mc *clientMockConn) SetWriteDeadline(_ time.Time) error { return nil } diff --git a/hotline/file_wrapper_test.go b/hotline/file_wrapper_test.go new file mode 100644 index 0000000..45fac94 --- /dev/null +++ b/hotline/file_wrapper_test.go @@ -0,0 +1,171 @@ +package hotline + +import ( + "errors" + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" +) + +func TestFile_DataFile(t *testing.T) { + t.Run("returns file info when data file exists", func(t *testing.T) { + mfs := &MockFileStore{} + mfi := &MockFileInfo{} + + f := &File{ + fs: mfs, + Name: "testfile.txt", + dataPath: "/files/testfile.txt", + incompletePath: "/files/testfile.txt.incomplete", + } + + mfs.On("Stat", "/files/testfile.txt").Return(mfi, nil) + + fi, err := f.DataFile() + + assert.NoError(t, err) + assert.Equal(t, mfi, fi) + mfs.AssertExpectations(t) + }) + + t.Run("returns file info from incomplete file when data file not found", func(t *testing.T) { + mfs := &MockFileStore{} + mfi := &MockFileInfo{} + + f := &File{ + fs: mfs, + Name: "testfile.txt", + dataPath: "/files/testfile.txt", + incompletePath: "/files/testfile.txt.incomplete", + } + + mfs.On("Stat", "/files/testfile.txt").Return(nil, os.ErrNotExist) + mfs.On("Stat", "/files/testfile.txt.incomplete").Return(mfi, nil) + + fi, err := f.DataFile() + + assert.NoError(t, err) + assert.Equal(t, mfi, fi) + mfs.AssertExpectations(t) + }) + + t.Run("returns error when neither data file nor incomplete file exists", func(t *testing.T) { + mfs := &MockFileStore{} + + f := &File{ + fs: mfs, + Name: "testfile.txt", + dataPath: "/files/testfile.txt", + incompletePath: "/files/testfile.txt.incomplete", + } + + mfs.On("Stat", "/files/testfile.txt").Return(nil, os.ErrNotExist) + mfs.On("Stat", "/files/testfile.txt.incomplete").Return(nil, os.ErrNotExist) + + fi, err := f.DataFile() + + assert.Nil(t, fi) + assert.EqualError(t, err, "file or directory not found") + mfs.AssertExpectations(t) + }) +} + +func TestFile_Move(t *testing.T) { + t.Run("succeeds when data rename works and meta files do not exist", func(t *testing.T) { + mfs := &MockFileStore{} + + f := &File{ + fs: mfs, + Name: "testfile.txt", + dataPath: "/files/testfile.txt", + incompletePath: "/files/testfile.txt.incomplete", + rsrcPath: "/files/.rsrc_testfile.txt", + infoPath: "/files/.info_testfile.txt", + } + + newPath := "/dest" + + mfs.On("Rename", "/files/testfile.txt", "/dest/testfile.txt").Return(nil) + mfs.On("Rename", "/files/testfile.txt.incomplete", "/dest/testfile.txt.incomplete").Return(os.ErrNotExist) + mfs.On("Rename", "/files/.rsrc_testfile.txt", "/dest/.rsrc_testfile.txt").Return(os.ErrNotExist) + mfs.On("Rename", "/files/.info_testfile.txt", "/dest/.info_testfile.txt").Return(os.ErrNotExist) + + err := f.Move(newPath) + + assert.NoError(t, err) + mfs.AssertExpectations(t) + }) + + t.Run("returns error when data file rename fails", func(t *testing.T) { + mfs := &MockFileStore{} + + f := &File{ + fs: mfs, + Name: "testfile.txt", + dataPath: "/files/testfile.txt", + incompletePath: "/files/testfile.txt.incomplete", + rsrcPath: "/files/.rsrc_testfile.txt", + infoPath: "/files/.info_testfile.txt", + } + + renameErr := errors.New("permission denied") + mfs.On("Rename", "/files/testfile.txt", "/dest/testfile.txt").Return(renameErr) + + err := f.Move("/dest") + + assert.ErrorIs(t, err, renameErr) + // Meta file renames should not have been called. + mfs.AssertNotCalled(t, "Rename", mock.Anything, "/dest/testfile.txt.incomplete") + mfs.AssertExpectations(t) + }) +} + +func TestFile_Delete(t *testing.T) { + t.Run("succeeds when RemoveAll works and meta files do not exist", func(t *testing.T) { + mfs := &MockFileStore{} + + f := &File{ + fs: mfs, + Name: "testfile.txt", + dataPath: "/files/testfile.txt", + incompletePath: "/files/testfile.txt.incomplete", + rsrcPath: "/files/.rsrc_testfile.txt", + infoPath: "/files/.info_testfile.txt", + } + + mfs.On("RemoveAll", "/files/testfile.txt").Return(nil) + mfs.On("Remove", "/files/testfile.txt.incomplete").Return(os.ErrNotExist) + mfs.On("Remove", "/files/.rsrc_testfile.txt").Return(os.ErrNotExist) + mfs.On("Remove", "/files/.info_testfile.txt").Return(os.ErrNotExist) + + err := f.Delete() + + assert.NoError(t, err) + mfs.AssertExpectations(t) + }) + + t.Run("returns error when RemoveAll fails", func(t *testing.T) { + mfs := &MockFileStore{} + + f := &File{ + fs: mfs, + Name: "testfile.txt", + dataPath: "/files/testfile.txt", + incompletePath: "/files/testfile.txt.incomplete", + rsrcPath: "/files/.rsrc_testfile.txt", + infoPath: "/files/.info_testfile.txt", + } + + removeErr := errors.New("permission denied") + mfs.On("RemoveAll", "/files/testfile.txt").Return(removeErr) + + err := f.Delete() + + assert.ErrorIs(t, err, removeErr) + // Meta file removes should not have been called. + mfs.AssertNotCalled(t, "Remove", mock.Anything) + mfs.AssertExpectations(t) + }) +} diff --git a/hotline/flattened_file_object_test.go b/hotline/flattened_file_object_test.go index c34af8a..6279157 100644 --- a/hotline/flattened_file_object_test.go +++ b/hotline/flattened_file_object_test.go @@ -1,6 +1,7 @@ package hotline import ( + "encoding/binary" "fmt" "testing" @@ -42,3 +43,144 @@ func TestFlatFileInformationFork_UnmarshalBinary(t *testing.T) { }) } } + +func TestFlatFileInformationFork_FriendlyType(t *testing.T) { + tests := []struct { + name string + typeSignature [4]byte + want []byte + }{ + { + name: "known type TEXT", + typeSignature: [4]byte{'T', 'E', 'X', 'T'}, + want: []byte("Text File"), + }, + { + name: "known type APPL", + typeSignature: [4]byte{'A', 'P', 'P', 'L'}, + want: []byte("Application Program"), + }, + { + name: "known type SIT!", + typeSignature: [4]byte{'S', 'I', 'T', '!'}, + want: []byte("StuffIt Archive"), + }, + { + name: "unknown type returns raw signature", + typeSignature: [4]byte{'J', 'P', 'E', 'G'}, + want: []byte("JPEG"), + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ffif := &FlatFileInformationFork{ + TypeSignature: tt.typeSignature, + } + assert.Equal(t, tt.want, ffif.FriendlyType()) + }) + } +} + +func TestFlatFileInformationFork_FriendlyCreator(t *testing.T) { + tests := []struct { + name string + creatorSignature [4]byte + want []byte + }{ + { + name: "known creator HTLC", + creatorSignature: [4]byte{'H', 'T', 'L', 'C'}, + want: []byte("Hotline"), + }, + { + name: "unknown creator returns raw signature", + creatorSignature: [4]byte{'o', 'g', 'l', 'e'}, + want: []byte("ogle"), + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ffif := &FlatFileInformationFork{ + CreatorSignature: tt.creatorSignature, + } + assert.Equal(t, tt.want, ffif.FriendlyCreator()) + }) + } +} + +func TestFlatFileInformationFork_SetComment(t *testing.T) { + tests := []struct { + name string + comment []byte + wantComment []byte + wantCommentSize [2]byte + }{ + { + name: "sets a short comment", + comment: []byte("hello"), + wantComment: []byte("hello"), + wantCommentSize: [2]byte{0x00, 0x05}, + }, + { + name: "sets an empty comment", + comment: []byte{}, + wantComment: []byte{}, + wantCommentSize: [2]byte{0x00, 0x00}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ffif := &FlatFileInformationFork{} + err := ffif.SetComment(tt.comment) + assert.NoError(t, err) + assert.Equal(t, tt.wantComment, ffif.Comment) + assert.Equal(t, tt.wantCommentSize, ffif.CommentSize) + }) + } +} + +func TestFlattenedFileObject_TransferSize(t *testing.T) { + tests := []struct { + name string + dataSize uint32 + resForkSize uint32 + offset int64 + wantNonZero bool + }{ + { + name: "calculates transfer size with zero offset", + dataSize: 100, + resForkSize: 50, + offset: 0, + wantNonZero: true, + }, + { + name: "calculates transfer size with non-zero offset", + dataSize: 200, + resForkSize: 100, + offset: 50, + wantNonZero: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ffo := &flattenedFileObject{} + + binary.BigEndian.PutUint32(ffo.FlatFileDataForkHeader.DataSize[:], tt.dataSize) + binary.BigEndian.PutUint32(ffo.FlatFileResForkHeader.DataSize[:], tt.resForkSize) + + result := ffo.TransferSize(tt.offset) + assert.Len(t, result, 4) + + size := binary.BigEndian.Uint32(result) + assert.Greater(t, size, uint32(0)) + + // With offset, the size should be smaller than without offset. + if tt.offset > 0 { + noOffsetResult := ffo.TransferSize(0) + noOffsetSize := binary.BigEndian.Uint32(noOffsetResult) + assert.Equal(t, noOffsetSize-uint32(tt.offset), size) + } + }) + } +} diff --git a/hotline/server_test.go b/hotline/server_test.go index abc2f9a..4bb3669 100644 --- a/hotline/server_test.go +++ b/hotline/server_test.go @@ -13,6 +13,7 @@ import ( "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" "golang.org/x/text/encoding" "golang.org/x/text/encoding/charmap" ) @@ -749,6 +750,107 @@ func TestServer_registerWithTrackers_EdgeCases(t *testing.T) { } } +func TestServer_CurrentStats(t *testing.T) { + stats := NewStats() + stats.Increment(StatCurrentlyConnected) + stats.Increment(StatDownloadCounter) + stats.Increment(StatDownloadCounter) + + srv := &Server{Stats: stats} + result := srv.CurrentStats() + + assert.Equal(t, 1, result["CurrentlyConnected"]) + assert.Equal(t, 2, result["DownloadCounter"]) + assert.Equal(t, 0, result["UploadsInProgress"]) +} + +func TestServer_sendTransaction(t *testing.T) { + t.Run("sends transaction to client connection", func(t *testing.T) { + wBuf := &bytes.Buffer{} + mockMgr := &MockClientMgr{} + mockMgr.On("Get", ClientID{0, 1}).Return(&ClientConn{ + Connection: &nopCloserRWC{Buffer: wBuf}, + }) + + srv := &Server{ClientMgr: mockMgr} + + tran := NewTransaction(TranChatMsg, ClientID{0, 1}, NewField(FieldData, []byte("hello"))) + err := srv.sendTransaction(tran) + + assert.NoError(t, err) + assert.Greater(t, wBuf.Len(), 0) + mockMgr.AssertExpectations(t) + }) + + t.Run("returns nil when client not found", func(t *testing.T) { + mockMgr := &MockClientMgr{} + mockMgr.On("Get", ClientID{0, 99}).Return((*ClientConn)(nil)) + + srv := &Server{ClientMgr: mockMgr} + + tran := NewTransaction(TranChatMsg, ClientID{0, 99}) + err := srv.sendTransaction(tran) + + assert.NoError(t, err) + mockMgr.AssertExpectations(t) + }) +} + +func TestServer_SendAll(t *testing.T) { + mockMgr := &MockClientMgr{} + mockMgr.On("List").Return([]*ClientConn{ + {ID: ClientID{0, 1}}, + {ID: ClientID{0, 2}}, + {ID: ClientID{0, 3}}, + }) + + outbox := make(chan Transaction, 10) + srv := &Server{ + ClientMgr: mockMgr, + outbox: outbox, + } + + srv.SendAll(TranChatMsg, NewField(FieldData, []byte("broadcast"))) + + assert.Len(t, outbox, 3) + + // Verify each transaction targets a different client + clientIDs := make(map[ClientID]bool) + for range 3 { + tran := <-outbox + clientIDs[tran.ClientID] = true + assert.Equal(t, TranChatMsg, tran.Type) + } + assert.True(t, clientIDs[ClientID{0, 1}]) + assert.True(t, clientIDs[ClientID{0, 2}]) + assert.True(t, clientIDs[ClientID{0, 3}]) + + mockMgr.AssertExpectations(t) +} + +type nopCloserRWC struct { + *bytes.Buffer +} + +func (n *nopCloserRWC) Close() error { return nil } + +func TestServer_NewClientConn(t *testing.T) { + mockMgr := &MockClientMgr{} + mockMgr.On("Add", mock.AnythingOfType("*hotline.ClientConn")).Return() + + srv := &Server{ClientMgr: mockMgr} + + rwc := &nopCloserRWC{Buffer: &bytes.Buffer{}} + + cc := srv.NewClientConn(rwc, "192.168.1.1:12345") + + assert.NotNil(t, cc) + assert.Equal(t, "192.168.1.1:12345", cc.RemoteAddr) + assert.Equal(t, []byte{0, 0}, cc.Icon) + assert.Equal(t, srv, cc.Server) + mockMgr.AssertExpectations(t) +} + func TestServer_registerWithAllTrackers(t *testing.T) { tests := []struct { name string @@ -819,3 +921,11 @@ func TestServer_registerWithAllTrackers(t *testing.T) { }) } } + +func TestSendBanMessage(t *testing.T) { + buf := &bytes.Buffer{} + sendBanMessage(buf, "You are banned") + + assert.Greater(t, buf.Len(), 0) + assert.Contains(t, buf.String(), "You are banned") +} -- cgit