diff options
| author | Jeff Halter <868228+jhalter@users.noreply.github.com> | 2026-07-10 11:18:37 -0700 |
|---|---|---|
| committer | Jeff Halter <868228+jhalter@users.noreply.github.com> | 2026-07-10 11:18:37 -0700 |
| commit | 38f66d6e2fc7f36a0c902565551d370e7af79182 (patch) | |
| tree | bd22db1e8cb97712a35d1e4a2efd02f562eea00e | |
| parent | 0802a90657ade00851e3d798fc80682438a091fa (diff) | |
Add tests for connection login sequence and client networking
Cover the two biggest untested plumbing paths in the hotline package:
- handleNewConnection (server login sequence): 3.2% -> 79.6%, via a
scripted in-memory ReadWriteCloser that feeds handshake + login bytes
and captures replies. Covers successful 1.5+ and legacy 1.2.3 flows,
no-agreement access, incorrect login, banned username, and banned IP.
- Client Connect/HandleTransactions/keepalive: 0% -> 83-93%, via a
loopback TCP server for Connect and preloaded mock conns for the
transaction loop and keepalive shutdown.
| -rw-r--r-- | hotline/client_connect_test.go | 178 | ||||
| -rw-r--r-- | hotline/handle_new_connection_test.go | 264 |
2 files changed, 442 insertions, 0 deletions
diff --git a/hotline/client_connect_test.go b/hotline/client_connect_test.go new file mode 100644 index 0000000..7c44929 --- /dev/null +++ b/hotline/client_connect_test.go @@ -0,0 +1,178 @@ +package hotline + +import ( + "bytes" + "context" + "errors" + "io" + "net" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// acceptHandshakeServer starts a TCP listener that, for a single connection, plays the server side +// of the handshake (optionally sending a bad response) and then reads the login transaction the +// client sends. The received login bytes are delivered on loginCh. It returns the listener's +// address so the client can dial it. +func acceptHandshakeServer(t *testing.T, goodHandshake bool) (addr string, loginCh <-chan []byte) { + t.Helper() + + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = ln.Close() }) + + ch := make(chan []byte, 1) + + go func() { + conn, err := ln.Accept() + if err != nil { + return + } + defer func() { _ = conn.Close() }() + + // Read the client handshake. + hs := make([]byte, handshakeSize) + if _, err := io.ReadFull(conn, hs); err != nil { + return + } + + if goodHandshake { + _, _ = conn.Write(ServerHandshake) + } else { + _, _ = conn.Write([]byte{0, 0, 0, 0, 0, 0, 0, 1}) + return + } + + // Read the login transaction that Connect sends next. + buf := make([]byte, 4096) + n, _ := conn.Read(buf) + ch <- buf[:n] + + // Hold the connection open briefly so the client's keepalive goroutine has a live socket. + time.Sleep(50 * time.Millisecond) + }() + + return ln.Addr().String(), ch +} + +func TestClient_Connect(t *testing.T) { + t.Run("completes handshake and sends login", func(t *testing.T) { + addr, loginCh := acceptHandshakeServer(t, true) + + c := newTestClient() + c.Pref.Username = "testuser" + + require.NoError(t, c.Connect(addr, "admin", "password")) + defer func() { _ = c.Disconnect() }() + + // Connect must establish the connection and arm the keepalive done channel. + assert.NotNil(t, c.Connection) + assert.NotNil(t, c.done) + + // The server should have received a well-formed login transaction carrying the + // obfuscated credentials and the username. + select { + case raw := <-loginCh: + var login Transaction + _, err := login.Write(raw) + require.NoError(t, err) + + assert.Equal(t, TranLogin, login.Type) + assert.Equal(t, "admin", login.GetField(FieldUserLogin).DecodeObfuscatedString()) + assert.Equal(t, "password", login.GetField(FieldUserPassword).DecodeObfuscatedString()) + assert.Equal(t, []byte("testuser"), login.GetField(FieldUserName).Data) + case <-time.After(2 * time.Second): + t.Fatal("server did not receive login transaction") + } + }) + + t.Run("returns error when dial fails", func(t *testing.T) { + // Reserve a port, then close the listener so nothing is accepting on it. + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + addr := ln.Addr().String() + require.NoError(t, ln.Close()) + + c := newTestClient() + err = c.Connect(addr, "admin", "password") + assert.Error(t, err) + }) + + t.Run("returns error on bad handshake response", func(t *testing.T) { + addr, _ := acceptHandshakeServer(t, false) + + c := newTestClient() + err := c.Connect(addr, "admin", "password") + assert.Error(t, err) + assert.Contains(t, err.Error(), "handshake") + }) +} + +func TestClient_HandleTransactions(t *testing.T) { + t.Run("dispatches queued transactions then reports termination", func(t *testing.T) { + // Two server-initiated transactions back to back, followed by EOF. + var buf bytes.Buffer + buf.Write(serializeTransaction(t, NewTransaction(TranChatMsg, [2]byte{0, 0}, NewField(FieldData, []byte("hi"))))) + buf.Write(serializeTransaction(t, NewTransaction(TranChatMsg, [2]byte{0, 0}, NewField(FieldData, []byte("there"))))) + + c := newTestClient() + c.Connection = &clientMockConn{RBuf: &buf, WBuf: &bytes.Buffer{}} + + var got []string + c.HandleFunc(TranChatMsg, func(_ context.Context, _ *Client, tr *Transaction) ([]Transaction, error) { + got = append(got, string(tr.GetField(FieldData).Data)) + return nil, nil + }) + + err := c.HandleTransactions(context.Background()) + + // At EOF the scanner stops and HandleTransactions reports the connection is gone. + assert.EqualError(t, err, "connection terminated") + assert.Equal(t, []string{"hi", "there"}, got) + }) + + t.Run("returns scanner error", func(t *testing.T) { + wantErr := errors.New("read boom") + c := newTestClient() + c.Connection = &errorConn{err: wantErr} + + err := c.HandleTransactions(context.Background()) + assert.ErrorIs(t, err, wantErr) + }) +} + +func TestClient_keepalive_stopsOnDone(t *testing.T) { + c := newTestClient() + done := make(chan struct{}) + + errCh := make(chan error, 1) + go func() { errCh <- c.keepalive(done) }() + + // Closing done must cause keepalive to return promptly (well before the 300s tick). + close(done) + + select { + case err := <-errCh: + assert.NoError(t, err) + case <-time.After(1 * time.Second): + t.Fatal("keepalive did not return after done was closed") + } +} + +// errorConn is a net.Conn whose Read always fails with a non-EOF error, so a bufio.Scanner over it +// surfaces the error via Scanner.Err(). +type errorConn struct { + err error +} + +func (e *errorConn) Read([]byte) (int, error) { return 0, e.err } +func (e *errorConn) Write(b []byte) (int, error) { return len(b), nil } +func (e *errorConn) Close() error { return nil } +func (e *errorConn) LocalAddr() net.Addr { return nil } +func (e *errorConn) RemoteAddr() net.Addr { return nil } +func (e *errorConn) SetDeadline(time.Time) error { return nil } +func (e *errorConn) SetReadDeadline(time.Time) error { return nil } +func (e *errorConn) SetWriteDeadline(time.Time) error { return nil } diff --git a/hotline/handle_new_connection_test.go b/hotline/handle_new_connection_test.go new file mode 100644 index 0000000..4a4930f --- /dev/null +++ b/hotline/handle_new_connection_test.go @@ -0,0 +1,264 @@ +package hotline + +import ( + "bytes" + "context" + "io" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/bcrypt" + "golang.org/x/text/encoding/charmap" +) + +// scriptedConn is an io.ReadWriteCloser that serves a fixed script of bytes to the server (the +// handshake plus a login transaction) and captures everything the server writes back. Reads are +// driven by handleNewConnection's main goroutine while writes come from the ClientConn.writeLoop +// goroutine, so the write buffer is mutex-protected for -race. +type scriptedConn struct { + in *bytes.Buffer + mu sync.Mutex + out bytes.Buffer +} + +func (c *scriptedConn) Read(p []byte) (int, error) { return c.in.Read(p) } + +func (c *scriptedConn) Write(p []byte) (int, error) { + c.mu.Lock() + defer c.mu.Unlock() + return c.out.Write(p) +} + +func (c *scriptedConn) Close() error { return nil } + +func (c *scriptedConn) written() []byte { + c.mu.Lock() + defer c.mu.Unlock() + return bytes.Clone(c.out.Bytes()) +} + +// serializeTransaction renders a transaction to its wire bytes. +func serializeTransaction(t *testing.T, tran Transaction) []byte { + t.Helper() + var buf bytes.Buffer + _, err := io.Copy(&buf, &tran) + require.NoError(t, err) + return buf.Bytes() +} + +// loginInput builds the byte script a client sends on connect: the handshake followed by a +// TranLogin transaction carrying the given credentials and optional extra fields. +func loginInput(t *testing.T, login, password string, extraFields ...Field) *bytes.Buffer { + t.Helper() + + fields := append([]Field{ + NewField(FieldUserLogin, EncodeString([]byte(login))), + NewField(FieldUserPassword, EncodeString([]byte(password))), + }, extraFields...) + + buf := bytes.NewBuffer(bytes.Clone(ClientHandshake)) + buf.Write(serializeTransaction(t, NewTransaction(TranLogin, [2]byte{0, 0}, fields...))) + return buf +} + +// newLoginTestServer builds a Server wired with the minimum dependencies handleNewConnection needs +// to run a login to completion. +func newLoginTestServer(accounts map[string]*Account) *Server { + return &Server{ + Config: Config{Name: "Test Server"}, + Logger: NewTestLogger(), + Stats: NewStats(), + ClientMgr: NewMemClientMgr(), + AccountManager: &mockAccountMgr{accounts: accounts}, + Agreement: bytes.NewReader([]byte("Welcome to the server")), + TextEncoder: charmap.Macintosh.NewEncoder(), + TextDecoder: charmap.Macintosh.NewDecoder(), + } +} + +// accountWithPassword returns an account whose stored password verifies against the given +// plaintext. The server authenticates against the obfuscated bytes carried on the wire (login +// never de-obfuscates FieldUserPassword), and account creation hashes those same obfuscated bytes, +// so the stored hash must cover EncodeString(password), not the plaintext. +func accountWithPassword(t *testing.T, login, password string, access ...int) *Account { + t.Helper() + hash, err := bcrypt.GenerateFromPassword(EncodeString([]byte(password)), bcrypt.MinCost) + require.NoError(t, err) + + acc := &Account{Login: login, Name: login, Password: string(hash)} + for _, a := range access { + acc.Access.Set(a) + } + return acc +} + +func TestServer_handleNewConnection_successfulLogin(t *testing.T) { + srv := newLoginTestServer(map[string]*Account{ + "admin": accountWithPassword(t, "admin", "password"), + }) + + conn := &scriptedConn{in: loginInput(t, "admin", "password", + NewField(FieldVersion, []byte{0, 0xc8}), + )} + + err := srv.handleNewConnection(context.Background(), conn, "192.168.1.5:1234") + require.NoError(t, err) + + // Wait for the writeLoop goroutine to flush the queued replies and exit. + srv.connWG.Wait() + + out := conn.written() + + // The server must complete the handshake first. + require.GreaterOrEqual(t, len(out), len(ServerHandshake)) + assert.Equal(t, ServerHandshake, out[:len(ServerHandshake)]) + + // The login reply, user-access, and agreement transactions should all have been sent, so the + // server name and agreement text land in the output stream. + assert.Contains(t, string(out), "Test Server") + assert.Contains(t, string(out), "Welcome to the server") + + // The connection should have been removed from the client manager on disconnect. + assert.Empty(t, srv.ClientMgr.List()) +} + +func TestServer_handleNewConnection_legacyLoginFlowNotifiesPresence(t *testing.T) { + srv := newLoginTestServer(map[string]*Account{ + // AccessAnyName lets the client keep the name it supplied at login. + "admin": accountWithPassword(t, "admin", "password", AccessAnyName), + }) + presence := &stubPresence{} + srv.Presence = presence + + // A username supplied in the login (and no version field) is the 1.2.3 login flow, which sets + // the nickname immediately and fires the presence/notify block. + conn := &scriptedConn{in: loginInput(t, "admin", "password", + NewField(FieldUserName, []byte("Administrator")), + )} + + require.NoError(t, srv.handleNewConnection(context.Background(), conn, "172.16.0.9:4000")) + srv.connWG.Wait() + + assert.Equal(t, []string{"admin"}, presence.connected) + assert.Equal(t, []string{"admin:Administrator"}, presence.renamed) + assert.Equal(t, []string{"admin:Administrator"}, presence.disconnected) +} + +func TestServer_handleNewConnection_noAgreementAccess(t *testing.T) { + srv := newLoginTestServer(map[string]*Account{ + "admin": accountWithPassword(t, "admin", "password", AccessNoAgreement), + }) + + conn := &scriptedConn{in: loginInput(t, "admin", "password", + NewField(FieldVersion, []byte{0, 0xc8}), + )} + + require.NoError(t, srv.handleNewConnection(context.Background(), conn, "10.0.0.1:5000")) + srv.connWG.Wait() + + // Accounts with AccessNoAgreement never receive the agreement body. + assert.NotContains(t, string(conn.written()), "Welcome to the server") +} + +func TestServer_handleNewConnection_incorrectLogin(t *testing.T) { + srv := newLoginTestServer(map[string]*Account{ + "admin": accountWithPassword(t, "admin", "password"), + }) + + conn := &scriptedConn{in: loginInput(t, "admin", "wrongpassword")} + + require.NoError(t, srv.handleNewConnection(context.Background(), conn, "192.168.1.5:1234")) + srv.connWG.Wait() + + assert.Contains(t, string(conn.written()), "Incorrect login") + assert.Empty(t, srv.ClientMgr.List()) +} + +func TestServer_handleNewConnection_bannedUsername(t *testing.T) { + srv := newLoginTestServer(map[string]*Account{ + "baduser": accountWithPassword(t, "baduser", "password"), + }) + ban := &stubBanMgr{usernameBanned: true} + srv.BanList = ban + + conn := &scriptedConn{in: loginInput(t, "baduser", "password")} + + require.NoError(t, srv.handleNewConnection(context.Background(), conn, "1.2.3.4:9999")) + srv.connWG.Wait() + + assert.Contains(t, string(conn.written()), "banned") + // The offending IP should have been added to the ban list. + assert.Equal(t, []string{"1.2.3.4"}, ban.added) + // A banned user never becomes a live connection. + assert.Empty(t, srv.ClientMgr.List()) +} + +func TestServer_handleNewConnection_bannedIP(t *testing.T) { + srv := newLoginTestServer(map[string]*Account{ + "admin": accountWithPassword(t, "admin", "password"), + }) + srv.BanList = &stubBanMgr{ipBanned: true} // permaban (nil expiry) + + conn := &scriptedConn{in: loginInput(t, "admin", "password")} + + require.NoError(t, srv.handleNewConnection(context.Background(), conn, "5.6.7.8:1111")) + srv.connWG.Wait() + + assert.Contains(t, string(conn.written()), "permanently banned") + assert.Empty(t, srv.ClientMgr.List()) +} + +func TestServer_handleNewConnection_handshakeFailure(t *testing.T) { + srv := newLoginTestServer(nil) + + // Wrong protocol bytes: handshake must fail before any login processing. + conn := &scriptedConn{in: bytes.NewBuffer([]byte("XXXXYYYY\x00\x01\x00\x02"))} + + err := srv.handleNewConnection(context.Background(), conn, "9.9.9.9:2222") + require.Error(t, err) + assert.Contains(t, err.Error(), "handshake") +} + +// stubPresence records PresenceTracker callbacks for assertion. +type stubPresence struct { + connected []string + renamed []string + disconnected []string +} + +func (s *stubPresence) UserConnected(login, ip string) { + s.connected = append(s.connected, login) +} +func (s *stubPresence) UserRenamed(login, oldNickname, newNickname, ip string) { + s.renamed = append(s.renamed, login+":"+newNickname) +} +func (s *stubPresence) UserDisconnected(login, nickname, ip string) { + s.disconnected = append(s.disconnected, login+":"+nickname) +} + +// stubBanMgr is a minimal BanMgr for exercising the ban branches of handleNewConnection. +type stubBanMgr struct { + usernameBanned bool + ipBanned bool + banUntil *time.Time + added []string +} + +func (s *stubBanMgr) Add(ip string, until *time.Time) error { + s.added = append(s.added, ip) + return nil +} +func (s *stubBanMgr) IsBanned(ip string) (bool, *time.Time) { return s.ipBanned, s.banUntil } +func (s *stubBanMgr) UnbanIP(ip string) error { return nil } +func (s *stubBanMgr) BanUsername(username string) error { return nil } +func (s *stubBanMgr) UnbanUsername(username string) error { return nil } +func (s *stubBanMgr) IsUsernameBanned(username string) bool { return s.usernameBanned } +func (s *stubBanMgr) BanNickname(nickname string) error { return nil } +func (s *stubBanMgr) UnbanNickname(nickname string) error { return nil } +func (s *stubBanMgr) IsNicknameBanned(nickname string) bool { return false } +func (s *stubBanMgr) ListBannedIPs() ([]string, error) { return nil, nil } +func (s *stubBanMgr) ListBannedUsernames() ([]string, error) { return nil, nil } +func (s *stubBanMgr) ListBannedNicknames() ([]string, error) { return nil, nil } |