aboutsummaryrefslogtreecommitdiff
path: root/hotline
diff options
context:
space:
mode:
Diffstat (limited to 'hotline')
-rw-r--r--hotline/client_conn.go140
-rw-r--r--hotline/client_conn_test.go57
-rw-r--r--hotline/server.go103
-rw-r--r--hotline/server_test.go42
4 files changed, 289 insertions, 53 deletions
diff --git a/hotline/client_conn.go b/hotline/client_conn.go
index 363376e..3419e0a 100644
--- a/hotline/client_conn.go
+++ b/hotline/client_conn.go
@@ -30,23 +30,26 @@ type ClientConn struct {
Connection io.ReadWriteCloser
RemoteAddr string
ID ClientID
- Icon []byte // TODO: make fixed size of 2
Version []byte // TODO: make fixed size of 2
- FlagsMU sync.Mutex // TODO: move into UserFlags struct
- Flags UserFlags
+ Account *Account
+ Server *Server // TODO: consider adding methods to interact with server
+ // The following fields hold mutable session state guarded by mu. They are read and written
+ // by multiple goroutines (the client's own transaction loop, other clients' handlers, and the
+ // server keepalive loop), so production code must use the accessor methods below. Direct
+ // field access is only safe before the connection is shared, e.g. in tests.
+ Flags UserFlags
UserName []byte
- Account *Account
+ Icon []byte // TODO: make fixed size of 2
IdleTime int
- Server *Server // TODO: consider adding methods to interact with server
AutoReply []byte
ClientFileTransferMgr ClientFileTransferMgr
Logger *slog.Logger
- mu sync.RWMutex
+ mu sync.RWMutex // guards the mutable session state fields above
sendCh chan Transaction
sendInit sync.Once
@@ -120,6 +123,109 @@ func (cc *ClientConn) writeLoop() {
func (cc *ClientConn) TextDecoder() *encoding.Decoder { return cc.Server.TextDecoder }
func (cc *ClientConn) TextEncoder() *encoding.Encoder { return cc.Server.TextEncoder }
+// SetFlag sets the user flag at position flag to v.
+func (cc *ClientConn) SetFlag(flag int, v uint) {
+ cc.mu.Lock()
+ defer cc.mu.Unlock()
+
+ cc.Flags.Set(flag, v)
+}
+
+// IsFlagSet reports whether the user flag at position flag is set.
+func (cc *ClientConn) IsFlagSet(flag int) bool {
+ cc.mu.RLock()
+ defer cc.mu.RUnlock()
+
+ return cc.Flags.IsSet(flag)
+}
+
+// FlagBytes returns a copy of the user flags bitmap, suitable for use as a transaction field.
+func (cc *ClientConn) FlagBytes() []byte {
+ cc.mu.RLock()
+ defer cc.mu.RUnlock()
+
+ flags := cc.Flags
+ return flags[:]
+}
+
+// SetUserName sets the client's display name.
+func (cc *ClientConn) SetUserName(name []byte) {
+ cc.mu.Lock()
+ defer cc.mu.Unlock()
+
+ cc.UserName = name
+}
+
+// GetUserName returns the client's display name. Callers must not modify the returned slice.
+func (cc *ClientConn) GetUserName() []byte {
+ cc.mu.RLock()
+ defer cc.mu.RUnlock()
+
+ return cc.UserName
+}
+
+// SetIcon sets the client's icon ID bytes.
+func (cc *ClientConn) SetIcon(icon []byte) {
+ cc.mu.Lock()
+ defer cc.mu.Unlock()
+
+ cc.Icon = icon
+}
+
+// GetIcon returns the client's icon ID bytes. Callers must not modify the returned slice.
+func (cc *ClientConn) GetIcon() []byte {
+ cc.mu.RLock()
+ defer cc.mu.RUnlock()
+
+ return cc.Icon
+}
+
+// SetAutoReply sets the client's away auto-reply message.
+func (cc *ClientConn) SetAutoReply(msg []byte) {
+ cc.mu.Lock()
+ defer cc.mu.Unlock()
+
+ cc.AutoReply = msg
+}
+
+// GetAutoReply returns the client's away auto-reply message. Callers must not modify the
+// returned slice.
+func (cc *ClientConn) GetAutoReply() []byte {
+ cc.mu.RLock()
+ defer cc.mu.RUnlock()
+
+ return cc.AutoReply
+}
+
+// incrementIdleTime adds interval seconds to the client's idle time. It returns true if this
+// crossed the idle threshold and marked the client as away, in which case the caller should
+// notify other clients of the change.
+func (cc *ClientConn) incrementIdleTime(interval int) bool {
+ cc.mu.Lock()
+ defer cc.mu.Unlock()
+
+ cc.IdleTime += interval
+ if cc.IdleTime > userIdleSeconds && !cc.Flags.IsSet(UserFlagAway) {
+ cc.Flags.Set(UserFlagAway, 1)
+ return true
+ }
+ return false
+}
+
+// clearIdleAndAway resets the client's idle timer. It returns true if the client was marked as
+// away and is no longer, in which case the caller should notify other clients of the change.
+func (cc *ClientConn) clearIdleAndAway() bool {
+ cc.mu.Lock()
+ defer cc.mu.Unlock()
+
+ cc.IdleTime = 0
+ if cc.Flags.IsSet(UserFlagAway) {
+ cc.Flags.Set(UserFlagAway, 0)
+ return true
+ }
+ return false
+}
+
func (cc *ClientConn) FileRoot() string {
if cc.Account.FileRoot != "" {
return cc.Account.FileRoot
@@ -201,23 +307,15 @@ func (cc *ClientConn) handleTransaction(transaction Transaction) {
}
if transaction.Type != TranKeepAlive {
- cc.mu.Lock()
- defer cc.mu.Unlock()
-
- // reset the user idle timer
- cc.IdleTime = 0
-
- // if user was previously idle, mark as not idle and notify other connected clients that
- // the user is no longer away
- if cc.Flags.IsSet(UserFlagAway) {
- cc.Flags.Set(UserFlagAway, 0)
-
+ // Reset the user idle timer. If the user was previously marked as away, notify other
+ // connected clients that the user is no longer away.
+ if cc.clearIdleAndAway() {
cc.SendAll(
TranNotifyChangeUser,
NewField(FieldUserID, cc.ID[:]),
- NewField(FieldUserFlags, cc.Flags[:]),
- NewField(FieldUserName, cc.UserName),
- NewField(FieldUserIconID, cc.Icon),
+ NewField(FieldUserFlags, cc.FlagBytes()),
+ NewField(FieldUserName, cc.GetUserName()),
+ NewField(FieldUserIconID, cc.GetIcon()),
)
}
}
@@ -328,7 +426,7 @@ func formatDownloadList(fts []FileTransfer) (s string) {
func (cc *ClientConn) String() string {
template := fmt.Sprintf(
userInfoTemplate,
- cc.UserName,
+ cc.GetUserName(),
cc.Account.Name,
cc.Account.Login,
cc.RemoteAddr,
diff --git a/hotline/client_conn_test.go b/hotline/client_conn_test.go
index 5c5463d..3312954 100644
--- a/hotline/client_conn_test.go
+++ b/hotline/client_conn_test.go
@@ -660,3 +660,60 @@ func TestClientConn_SendDisconnectRace(t *testing.T) {
wg.Wait()
}
+
+func TestClientConn_incrementIdleTime(t *testing.T) {
+ cc := &ClientConn{}
+
+ // Increment until just below the idle threshold: not yet away.
+ for i := 0; i < userIdleSeconds/idleCheckInterval; i++ {
+ assert.False(t, cc.incrementIdleTime(idleCheckInterval))
+ }
+ assert.False(t, cc.IsFlagSet(UserFlagAway))
+
+ // The increment that crosses the threshold marks the client away exactly once.
+ assert.True(t, cc.incrementIdleTime(idleCheckInterval))
+ assert.True(t, cc.IsFlagSet(UserFlagAway))
+ assert.False(t, cc.incrementIdleTime(idleCheckInterval), "already-away client should not be marked away again")
+}
+
+func TestClientConn_clearIdleAndAway(t *testing.T) {
+ cc := &ClientConn{IdleTime: 500}
+
+ // Not away: idle timer resets, no notification needed.
+ assert.False(t, cc.clearIdleAndAway())
+ assert.Equal(t, 0, cc.IdleTime)
+
+ // Away: flag clears and the caller is told to notify.
+ cc.SetFlag(UserFlagAway, 1)
+ assert.True(t, cc.clearIdleAndAway())
+ assert.False(t, cc.IsFlagSet(UserFlagAway))
+ assert.False(t, cc.clearIdleAndAway(), "second clear should report no change")
+}
+
+// TestClientConn_sessionStateRace exercises concurrent access to the mutable session state through
+// the accessor methods. Run with -race.
+func TestClientConn_sessionStateRace(t *testing.T) {
+ cc := &ClientConn{}
+
+ var wg sync.WaitGroup
+ for range 4 {
+ wg.Add(1)
+ go func() {
+ defer wg.Done()
+ for i := range 100 {
+ cc.SetUserName(fmt.Appendf(nil, "user-%d", i))
+ _ = cc.GetUserName()
+ cc.SetIcon([]byte{0, byte(i)})
+ _ = cc.GetIcon()
+ cc.SetAutoReply([]byte("brb"))
+ _ = cc.GetAutoReply()
+ cc.SetFlag(UserFlagRefusePM, uint(i%2))
+ _ = cc.IsFlagSet(UserFlagRefusePM)
+ _ = cc.FlagBytes()
+ _ = cc.incrementIdleTime(idleCheckInterval)
+ _ = cc.clearIdleAndAway()
+ }
+ }()
+ }
+ wg.Wait()
+}
diff --git a/hotline/server.go b/hotline/server.go
index 3459346..8a530c2 100644
--- a/hotline/server.go
+++ b/hotline/server.go
@@ -34,7 +34,7 @@ type Server struct {
NetInterface string
Port int
- rateLimiters map[string]*rate.Limiter
+ rateLimiters map[string]*rateLimiterEntry
rateLimitersMu sync.Mutex
handlers map[TranType]HandlerFunc
@@ -49,7 +49,9 @@ type Server struct {
FS FileStore // Storage backend to use for File storage
Agreement io.ReadSeeker
- Banner []byte
+
+ banner []byte // server banner image; guarded by bannerMu as it is replaced on config reload
+ bannerMu sync.RWMutex
FileTransferMgr FileTransferMgr
ChatMgr ChatManager
@@ -80,6 +82,22 @@ func (s *Server) initShutdownCh() {
s.shutdownInit.Do(func() { s.shutdownCh = make(chan struct{}) })
}
+// Banner returns the server banner image. Callers must not modify the returned slice.
+func (s *Server) Banner() []byte {
+ s.bannerMu.RLock()
+ defer s.bannerMu.RUnlock()
+
+ return s.banner
+}
+
+// SetBanner replaces the server banner image.
+func (s *Server) SetBanner(banner []byte) {
+ s.bannerMu.Lock()
+ defer s.bannerMu.Unlock()
+
+ s.banner = banner
+}
+
type Option = func(s *Server)
func WithConfig(config Config) func(s *Server) {
@@ -129,7 +147,7 @@ type ServerConfig struct {
func NewServer(options ...Option) (*Server, error) {
server := Server{
handlers: make(map[TranType]HandlerFunc),
- rateLimiters: make(map[string]*rate.Limiter),
+ rateLimiters: make(map[string]*rateLimiterEntry),
FS: &OSFileStore{},
ChatMgr: NewMemChatManager(),
ClientMgr: NewMemClientMgr(),
@@ -274,6 +292,29 @@ func (s *Server) Send(t Transaction) {
// 0.5 = 1 connection every 2 seconds
const perIPRateLimit = rate.Limit(0.5)
+// rateLimiterTTL is how long the rate limiter for an idle IP address is retained before eviction.
+const rateLimiterTTL = 7 * 24 * time.Hour
+
+type rateLimiterEntry struct {
+ limiter *rate.Limiter
+ lastSeen time.Time
+}
+
+// sweepRateLimiters evicts rate limiters for IP addresses that have not connected recently, so
+// the rate limiter map does not grow unbounded over the lifetime of the server.
+func (s *Server) sweepRateLimiters() {
+ cutoff := time.Now().Add(-rateLimiterTTL)
+
+ s.rateLimitersMu.Lock()
+ defer s.rateLimitersMu.Unlock()
+
+ for ip, entry := range s.rateLimiters {
+ if entry.lastSeen.Before(cutoff) {
+ delete(s.rateLimiters, ip)
+ }
+ }
+}
+
func (s *Server) Serve(ctx context.Context, ln net.Listener) error {
for {
conn, err := ln.Accept()
@@ -302,11 +343,13 @@ func (s *Server) Serve(ctx context.Context, ln net.Listener) error {
// Check if we have an existing rate limit for the IP and create one if we do not.
s.rateLimitersMu.Lock()
- rl, ok := s.rateLimiters[ipAddr]
+ entry, ok := s.rateLimiters[ipAddr]
if !ok {
- rl = rate.NewLimiter(perIPRateLimit, 1)
- s.rateLimiters[ipAddr] = rl
+ entry = &rateLimiterEntry{limiter: rate.NewLimiter(perIPRateLimit, 1)}
+ s.rateLimiters[ipAddr] = entry
}
+ entry.lastSeen = time.Now()
+ rl := entry.limiter
s.rateLimitersMu.Unlock()
// Check if the rate limit is exceeded and close the connection if so.
@@ -417,6 +460,7 @@ const (
// keepaliveHandler runs every idleCheckInterval seconds and increments a user's idle time by idleCheckInterval seconds.
// If the updated idle time exceeds userIdleSeconds and the user was not previously idle, we notify all connected clients
// that the user has gone idle. For most clients, this turns the user grey in the user list.
+// It also sweeps stale per-IP rate limiters on each tick.
func (s *Server) keepaliveHandler(ctx context.Context) {
ticker := time.NewTicker(idleCheckInterval * time.Second)
defer ticker.Stop()
@@ -427,23 +471,18 @@ func (s *Server) keepaliveHandler(ctx context.Context) {
return
case <-ticker.C:
for _, c := range s.ClientMgr.List() {
- c.mu.Lock()
- c.IdleTime += idleCheckInterval
-
- // Check if the user
- if c.IdleTime > userIdleSeconds && !c.Flags.IsSet(UserFlagAway) {
- c.Flags.Set(UserFlagAway, 1)
-
+ if c.incrementIdleTime(idleCheckInterval) {
c.SendAll(
TranNotifyChangeUser,
NewField(FieldUserID, c.ID[:]),
- NewField(FieldUserFlags, c.Flags[:]),
- NewField(FieldUserName, c.UserName),
- NewField(FieldUserIconID, c.Icon),
+ NewField(FieldUserFlags, c.FlagBytes()),
+ NewField(FieldUserName, c.GetUserName()),
+ NewField(FieldUserIconID, c.GetIcon()),
)
}
- c.mu.Unlock()
}
+
+ s.sweepRateLimiters()
}
}
}
@@ -547,8 +586,8 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser
defer func() {
if s.Redis != nil {
s.Redis.SRem(context.Background(), RedisKeyOnline, login+"::"+ipAddr)
- if len(c.UserName) != 0 {
- s.Redis.SRem(context.Background(), RedisKeyOnline, login+":"+string(c.UserName)+":"+ipAddr)
+ if userName := c.GetUserName(); len(userName) != 0 {
+ s.Redis.SRem(context.Background(), RedisKeyOnline, login+":"+string(userName)+":"+ipAddr)
}
}
c.Disconnect()
@@ -574,7 +613,7 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser
}
if clientLogin.GetField(FieldUserIconID).Data != nil {
- c.Icon = clientLogin.GetField(FieldUserIconID).Data
+ c.SetIcon(clientLogin.GetField(FieldUserIconID).Data)
}
c.Account = c.Server.AccountManager.Get(login)
@@ -584,14 +623,14 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser
if clientLogin.GetField(FieldUserName).Data != nil {
if c.Authorize(AccessAnyName) {
- c.UserName = clientLogin.GetField(FieldUserName).Data
+ c.SetUserName(clientLogin.GetField(FieldUserName).Data)
} else {
- c.UserName = []byte(c.Account.Name)
+ c.SetUserName([]byte(c.Account.Name))
}
}
if c.Authorize(AccessDisconUser) {
- c.Flags.Set(UserFlagAdmin, 1)
+ c.SetFlag(UserFlagAdmin, 1)
}
c.Send(c.NewReply(&clientLogin,
@@ -620,18 +659,18 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser
// If the client has provided a username as part of the login, we can infer that it is using the 1.2.3 login
// flow and not the 1.5+ flow.
- if len(c.UserName) != 0 {
+ if userName := c.GetUserName(); len(userName) != 0 {
// Add the client username to the logger. For 1.5+ clients, we don't have this information yet as it comes as
// part of TranAgreed
- c.Logger = c.Logger.With("name", string(c.UserName))
+ c.Logger = c.Logger.With("name", string(userName))
c.Logger.Info("Login successful")
// Update the Redis set with the new information
- if s.Redis != nil && len(c.UserName) != 0 {
+ if s.Redis != nil {
// Remove old entry (login::ip)
s.Redis.SRem(context.Background(), RedisKeyOnline, login+"::"+ipAddr)
// Add new entry with login, nickname, ip
- s.Redis.SAdd(context.Background(), RedisKeyOnline, login+":"+string(c.UserName)+":"+ipAddr)
+ s.Redis.SAdd(context.Background(), RedisKeyOnline, login+":"+string(userName)+":"+ipAddr)
}
// Notify other clients on the server that the new user has logged in. For 1.5+ clients we don't have this
@@ -639,10 +678,10 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser
for _, t := range c.NotifyOthers(
NewTransaction(
TranNotifyChangeUser, [2]byte{0, 0},
- NewField(FieldUserName, c.UserName),
+ NewField(FieldUserName, userName),
NewField(FieldUserID, c.ID[:]),
- NewField(FieldUserIconID, c.Icon),
- NewField(FieldUserFlags, c.Flags[:]),
+ NewField(FieldUserIconID, c.GetIcon()),
+ NewField(FieldUserFlags, c.FlagBytes()),
),
) {
c.Server.Send(t)
@@ -701,7 +740,7 @@ func (s *Server) handleFileTransfer(ctx context.Context, rwc io.ReadWriter) erro
rLogger := s.Logger.With(
"remoteAddr", remoteAddr,
"login", fileTransfer.ClientConn.Account.Login,
- "Name", string(fileTransfer.ClientConn.UserName),
+ "Name", string(fileTransfer.ClientConn.GetUserName()),
)
fullPath, err := ReadPath(fileTransfer.FileRoot, fileTransfer.FilePath, fileTransfer.FileName, s.TextDecoder)
@@ -711,7 +750,7 @@ func (s *Server) handleFileTransfer(ctx context.Context, rwc io.ReadWriter) erro
switch fileTransfer.Type {
case BannerDownload:
- if _, err := io.Copy(rwc, bytes.NewBuffer(s.Banner)); err != nil {
+ if _, err := io.Copy(rwc, bytes.NewReader(s.Banner())); err != nil {
return fmt.Errorf("banner download: %w", err)
}
case FileDownload:
diff --git a/hotline/server_test.go b/hotline/server_test.go
index a6c5e19..ebcd073 100644
--- a/hotline/server_test.go
+++ b/hotline/server_test.go
@@ -10,6 +10,7 @@ import (
"net"
"os"
"strings"
+ "sync"
"testing"
"time"
@@ -18,6 +19,7 @@ import (
"github.com/stretchr/testify/require"
"golang.org/x/text/encoding"
"golang.org/x/text/encoding/charmap"
+ "golang.org/x/time/rate"
)
type mockReadWriter struct {
@@ -1036,3 +1038,43 @@ func TestServer_Shutdown_stopsListenAndServe(t *testing.T) {
t.Fatal("ListenAndServe did not return after Shutdown")
}
}
+
+func TestServer_Banner(t *testing.T) {
+ srv := &Server{}
+ assert.Nil(t, srv.Banner())
+
+ srv.SetBanner([]byte("banner-data"))
+ assert.Equal(t, []byte("banner-data"), srv.Banner())
+
+ // Concurrent reads and writes; run with -race.
+ var wg sync.WaitGroup
+ for range 4 {
+ wg.Add(1)
+ go func() {
+ defer wg.Done()
+ for i := range 100 {
+ srv.SetBanner([]byte{byte(i)})
+ _ = srv.Banner()
+ }
+ }()
+ }
+ wg.Wait()
+}
+
+func TestServer_sweepRateLimiters(t *testing.T) {
+ srv := &Server{rateLimiters: make(map[string]*rateLimiterEntry)}
+
+ srv.rateLimiters["10.0.0.1"] = &rateLimiterEntry{
+ limiter: rate.NewLimiter(perIPRateLimit, 1),
+ lastSeen: time.Now().Add(-rateLimiterTTL - time.Minute), // stale
+ }
+ srv.rateLimiters["10.0.0.2"] = &rateLimiterEntry{
+ limiter: rate.NewLimiter(perIPRateLimit, 1),
+ lastSeen: time.Now(), // fresh
+ }
+
+ srv.sweepRateLimiters()
+
+ assert.NotContains(t, srv.rateLimiters, "10.0.0.1", "stale rate limiter should be evicted")
+ assert.Contains(t, srv.rateLimiters, "10.0.0.2", "fresh rate limiter should be retained")
+}