aboutsummaryrefslogtreecommitdiff
path: root/hotline/server.go
diff options
context:
space:
mode:
Diffstat (limited to 'hotline/server.go')
-rw-r--r--hotline/server.go103
1 files changed, 71 insertions, 32 deletions
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: