diff options
Diffstat (limited to 'hotline')
| -rw-r--r-- | hotline/client_conn.go | 140 | ||||
| -rw-r--r-- | hotline/client_conn_test.go | 57 | ||||
| -rw-r--r-- | hotline/server.go | 103 | ||||
| -rw-r--r-- | hotline/server_test.go | 42 |
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") +} |